From 09d2ca7abffb8205fbc55a8cbfaef5bdb222eca1 Mon Sep 17 00:00:00 2001 From: RainbowBird Date: Sat, 28 Jun 2025 13:34:17 +0800 Subject: [PATCH] feat: progress emitter (#237) --- apps/stage-tamagotchi/src-tauri/src/lib.rs | 21 +- .../src-tauri/src/whisper/mod.rs | 276 +----------------- .../src-tauri/src/whisper/progress.rs | 93 ++++++ .../src-tauri/src/whisper/whisper.rs | 247 ++++++++++++++++ apps/stage-tamagotchi/src/pages/index.vue | 5 +- 5 files changed, 360 insertions(+), 282 deletions(-) create mode 100644 apps/stage-tamagotchi/src-tauri/src/whisper/progress.rs create mode 100644 apps/stage-tamagotchi/src-tauri/src/whisper/whisper.rs diff --git a/apps/stage-tamagotchi/src-tauri/src/lib.rs b/apps/stage-tamagotchi/src-tauri/src/lib.rs index 00c3e10f7..709bb0a6d 100644 --- a/apps/stage-tamagotchi/src-tauri/src/lib.rs +++ b/apps/stage-tamagotchi/src-tauri/src/lib.rs @@ -97,7 +97,9 @@ async fn stop_click_through(window: tauri::Window) -> Result<(), String> { Ok(()) } -fn load_whisper_model() -> Result { +fn load_whisper_model(window: tauri::Window) -> anyhow::Result<()> { + let progress_manager = crate::whisper::progress::ModelLoadProgressEmitterManager::new(window); + // Determine device to use let device = if candle_core::utils::cuda_is_available() { candle_core::Device::new_cuda(0)? @@ -109,12 +111,19 @@ fn load_whisper_model() -> Result { info!("Using device: {device:?}"); - let whisper_model = whisper::WhichWhisperModel::Tiny; + let whisper_model = whisper::whisper::WhichWhisperModel::Tiny; info!("Loading whisper model: {:?}", whisper_model); - let model = whisper::WhisperProcessor::new(whisper_model, device.clone())?; - Ok(model) + let _ = whisper::whisper::WhisperProcessor::new(whisper_model, device.clone(), progress_manager)?; + Ok(()) +} + +#[tauri::command] +async fn load_models(window: tauri::Window) -> Result<(), String> { + println!("load_models"); + load_whisper_model(window).unwrap(); + Ok(()) } #[cfg_attr(mobile, tauri::mobile_entry_point)] @@ -162,9 +171,6 @@ pub fn run() { )?; } - // Load whisper model - load_whisper_model()?; - // TODO: i18n let quit_item = MenuItem::with_id(app, "quit", "Quit", true, None::<&str>)?; let settings_item = MenuItem::with_id(app, "settings", "Settings", true, None::<&str>)?; @@ -231,6 +237,7 @@ pub fn run() { stop_monitor, start_click_through, stop_click_through, + load_models, ]) .build(tauri::generate_context!()) .expect("error while building tauri application") diff --git a/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs b/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs index 105252993..07f130391 100644 --- a/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs +++ b/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs @@ -1,274 +1,2 @@ -// https://github.com/proj-airi/candle-examples/blob/3a6783a4788333e7a117f4f8548f02b2fbfec2ed/apps/silero-vad-whisper-realtime/src/whisper.rs -use anyhow::Result; -use byteorder::{ByteOrder, LittleEndian}; -use candle_core::{Device, IndexOp, Tensor}; -use candle_nn::VarBuilder; -use candle_transformers::models::whisper::{self as whisper_model, Config, audio}; -use clap::ValueEnum; -use hf_hub::{ - Repo, - RepoType, - api::sync::{Api, ApiBuilder}, -}; -use tokenizers::Tokenizer; - -pub enum WhisperModel { - Normal(whisper_model::model::Whisper), -} - -impl WhisperModel { - pub fn encoder_forward( - &mut self, - x: &Tensor, - flush: bool, - ) -> candle_core::Result { - match self { - Self::Normal(model) => model.encoder.forward(x, flush), - } - } - - pub fn decoder_forward( - &mut self, - x: &Tensor, - encoder_out: &Tensor, - flush: bool, - ) -> candle_core::Result { - match self { - Self::Normal(model) => model.decoder.forward(x, encoder_out, flush), - } - } - - pub fn decoder_final_linear( - &self, - x: &Tensor, - ) -> candle_core::Result { - match self { - Self::Normal(model) => model.decoder.final_linear(x), - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] -pub enum WhichWhisperModel { - Tiny, - #[value(name = "tiny.en")] - TinyEn, - Base, - #[value(name = "base.en")] - BaseEn, - Small, - #[value(name = "small.en")] - SmallEn, - Medium, - #[value(name = "medium.en")] - MediumEn, - Large, - LargeV2, - LargeV3, - LargeV3Turbo, - #[value(name = "distil-medium.en")] - DistilMediumEn, - #[value(name = "distil-large-v2")] - DistilLargeV2, -} - -impl WhichWhisperModel { - pub const fn model_and_revision(self) -> (&'static str, &'static str) { - match self { - Self::Tiny => ("openai/whisper-tiny", "main"), - Self::TinyEn => ("openai/whisper-tiny.en", "refs/pr/15"), - Self::Base => ("openai/whisper-base", "refs/pr/22"), - Self::BaseEn => ("openai/whisper-base.en", "refs/pr/13"), - Self::Small => ("openai/whisper-small", "main"), - Self::SmallEn => ("openai/whisper-small.en", "refs/pr/10"), - Self::Medium => ("openai/whisper-medium", "main"), - Self::MediumEn => ("openai/whisper-medium.en", "main"), - Self::Large => ("openai/whisper-large", "refs/pr/36"), - Self::LargeV2 => ("openai/whisper-large-v2", "refs/pr/57"), - Self::LargeV3 => ("openai/whisper-large-v3", "main"), - Self::LargeV3Turbo => ("openai/whisper-large-v3-turbo", "main"), - Self::DistilMediumEn => ("distil-whisper/distil-medium.en", "main"), - Self::DistilLargeV2 => ("distil-whisper/distil-large-v2", "main"), - } - } -} - -pub struct WhisperProcessor { - pub model: WhisperModel, - pub tokenizer: Tokenizer, - pub config: Config, - pub mel_filters: Vec, - pub device: Device, -} - -pub struct WhisperProgressEmitter { - filename: String, -} - -impl WhisperProgressEmitter { - fn new(filename: &str) -> Self { - Self { - filename: filename.to_string(), - } - } -} - -impl hf_hub::api::Progress for WhisperProgressEmitter { - fn init( - &mut self, - size: usize, - filename: &str, - ) { - todo!(); - } - - fn update( - &mut self, - size: usize, - ) { - todo!(); - } - - fn finish(&mut self) { - todo!(); - } -} - -impl WhisperProcessor { - pub fn new( - model: WhichWhisperModel, - device: Device, - ) -> Result { - // Load the Whisper model based on the provided model type - let api = ApiBuilder::new().with_progress(false).build()?; - let (model_id, revision) = model.model_and_revision(); - let repo = api.repo(Repo::with_revision( - model_id.to_string(), - RepoType::Model, - revision.to_string(), - )); - - let config_filename = - repo.download_with_progress("config.json", WhisperProgressEmitter::new("config.json"))?; - // TODO: for RainbowBird - let tokenizer_filename = repo.get("tokenizer.json")?; - // TODO: for RainbowBird - let model_filename = repo.get("model.safetensors")?; - - let config: Config = serde_json::from_str(&std::fs::read_to_string(config_filename)?)?; - let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(anyhow::Error::msg)?; - - println!("Loading Whisper model from: {:?}", model_filename.display()); - - // SAFETY: This is safe because we are using a mmaped file and the safetensors library guarantees that the data is valid. - let var_builder = unsafe { - VarBuilder::from_mmaped_safetensors(&[model_filename], whisper_model::DTYPE, &device)? - }; - - let model = WhisperModel::Normal(whisper_model::model::Whisper::load( - &var_builder, - config.clone(), - )?); - - let mel_bytes = match config.num_mel_bins { - 80 => include_bytes!("./melfilters.bytes").as_slice(), - 128 => include_bytes!("./melfilters128.bytes").as_slice(), - num_mel_bins => anyhow::bail!("Unsupported number of mel bins: {}", num_mel_bins), - }; - - let mut mel_filters = vec![0f32; mel_bytes.len() / 4]; - ::read_f32_into(mel_bytes, &mut mel_filters); - - Ok(Self { - model, - tokenizer, - config, - mel_filters, - device, - }) - } - - pub fn transcribe( - &mut self, - audio: &[f32], - ) -> Result { - // Convert PCM to mel spectrogram - let mel = audio::pcm_to_mel(&self.config, audio, &self.mel_filters); - let mel_len = mel.len(); - let mel = Tensor::from_vec( - mel, - ( - 1, - self.config.num_mel_bins, - mel_len / self.config.num_mel_bins, - ), - &self.device, - )?; - - // Run inference - let audio_features = self.model.encoder_forward(&mel, true)?; - // Simple greedy decoding - let tokens = self.decode_greedy(&audio_features)?; - - let text = self - .tokenizer - .decode(&tokens, true) - .map_err(anyhow::Error::msg)?; - - Ok(text) - } - - fn decode_greedy( - &mut self, - audio_features: &Tensor, - ) -> Result> { - let mut tokens = vec![ - self.token_id(whisper_model::SOT_TOKEN)?, - self.token_id(whisper_model::TRANSCRIBE_TOKEN)?, - self.token_id(whisper_model::NO_TIMESTAMPS_TOKEN)?, - ]; - - let max_len = 50; // Short sequence for real-time processing - - for i in 0..max_len { - let tokens_t = Tensor::new(tokens.as_slice(), &self.device)?.unsqueeze(0)?; - let ys = self - .model - .decoder_forward(&tokens_t, audio_features, i == 0)?; - - let (_, seq_len, _) = ys.dims3()?; - let logits = self - .model - .decoder_final_linear(&ys.i((..1, seq_len - 1..))?)? - .i(0)? - .i(0)?; - - // Get most likely token - let logits_v: Vec = logits.to_vec1()?; - let next_token = logits_v - .iter() - .enumerate() - .max_by(|(_, a), (_, b)| a.total_cmp(b)) - .map(|(i, _)| u32::try_from(i).unwrap()) - .unwrap(); - - if next_token == self.token_id(whisper_model::EOT_TOKEN)? { - break; - } - - tokens.push(next_token); - } - - Ok(tokens) - } - - fn token_id( - &self, - token: &str, - ) -> Result { - self - .tokenizer - .token_to_id(token) - .ok_or_else(|| anyhow::anyhow!("Token not found: {}", token)) - } -} +pub mod progress; +pub mod whisper; diff --git a/apps/stage-tamagotchi/src-tauri/src/whisper/progress.rs b/apps/stage-tamagotchi/src-tauri/src/whisper/progress.rs new file mode 100644 index 000000000..595421665 --- /dev/null +++ b/apps/stage-tamagotchi/src-tauri/src/whisper/progress.rs @@ -0,0 +1,93 @@ +use log::error; +use tauri::Emitter; + +#[derive(Clone)] +pub struct ModelLoadProgressEmitterManager { + window: tauri::Window, +} + +impl ModelLoadProgressEmitterManager { + pub fn new(window: tauri::Window) -> Self { + Self { window } + } + + pub fn new_processor( + self, + filename: &str, + ) -> ModelLoadProgressEmitter { + ModelLoadProgressEmitter::new(self.window, filename) + } +} + +pub struct ModelLoadProgressEmitter { + filename: String, + total_size: usize, + progress: f32, + window: tauri::Window, +} + +impl ModelLoadProgressEmitter { + fn new( + window: tauri::Window, + filename: &str, + ) -> Self { + Self { + filename: filename.to_string(), + total_size: 0, + progress: 0.0, + window, + } + } +} + +impl hf_hub::api::Progress for ModelLoadProgressEmitter { + fn init( + &mut self, + size: usize, + _: &str, + ) { + self.total_size = size; + self.progress = 0.0; + self + .window + .emit( + "tauri-app:model-load-progress", + (self.filename.clone(), self.progress), + ) + .map_err(|err| { + error!("Failed to emit model-load-progress: {:?}", err); + }) + .unwrap(); + } + + fn update( + &mut self, + size: usize, + ) { + self.progress = size as f32 / self.total_size as f32 * 100.0; + self + .window + .emit( + "tauri-app:model-load-progress", + (self.filename.clone(), self.progress), + ) + .map_err(|err| { + error!("Failed to emit model-load-progress: {:?}", err); + }) + .unwrap(); + } + + fn finish(&mut self) { + self.progress = 100.0; + self + .window + .emit( + "tauri-app:model-load-progress", + (self.filename.clone(), self.progress), + ) + .map_err(|err| { + error!("Failed to emit model-load-progress: {:?}", err); + }) + .unwrap(); + } +} diff --git a/apps/stage-tamagotchi/src-tauri/src/whisper/whisper.rs b/apps/stage-tamagotchi/src-tauri/src/whisper/whisper.rs new file mode 100644 index 000000000..67cf9bee9 --- /dev/null +++ b/apps/stage-tamagotchi/src-tauri/src/whisper/whisper.rs @@ -0,0 +1,247 @@ +// https://github.com/proj-airi/candle-examples/blob/3a6783a4788333e7a117f4f8548f02b2fbfec2ed/apps/silero-vad-whisper-realtime/src/whisper.rs +use anyhow::Result; +use byteorder::{ByteOrder, LittleEndian}; +use candle_core::{Device, IndexOp, Tensor}; +use candle_nn::VarBuilder; +use candle_transformers::models::whisper::{self as whisper_model, Config, audio}; +use clap::ValueEnum; +use hf_hub::{Repo, RepoType, api::sync::ApiBuilder}; +use tokenizers::Tokenizer; + +use crate::whisper::progress; + +pub enum WhisperModel { + Normal(whisper_model::model::Whisper), +} + +impl WhisperModel { + pub fn encoder_forward( + &mut self, + x: &Tensor, + flush: bool, + ) -> candle_core::Result { + match self { + Self::Normal(model) => model.encoder.forward(x, flush), + } + } + + pub fn decoder_forward( + &mut self, + x: &Tensor, + encoder_out: &Tensor, + flush: bool, + ) -> candle_core::Result { + match self { + Self::Normal(model) => model.decoder.forward(x, encoder_out, flush), + } + } + + pub fn decoder_final_linear( + &self, + x: &Tensor, + ) -> candle_core::Result { + match self { + Self::Normal(model) => model.decoder.final_linear(x), + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)] +pub enum WhichWhisperModel { + Tiny, + #[value(name = "tiny.en")] + TinyEn, + Base, + #[value(name = "base.en")] + BaseEn, + Small, + #[value(name = "small.en")] + SmallEn, + Medium, + #[value(name = "medium.en")] + MediumEn, + Large, + LargeV2, + LargeV3, + LargeV3Turbo, + #[value(name = "distil-medium.en")] + DistilMediumEn, + #[value(name = "distil-large-v2")] + DistilLargeV2, +} + +impl WhichWhisperModel { + pub const fn model_and_revision(self) -> (&'static str, &'static str) { + match self { + Self::Tiny => ("openai/whisper-tiny", "main"), + Self::TinyEn => ("openai/whisper-tiny.en", "refs/pr/15"), + Self::Base => ("openai/whisper-base", "refs/pr/22"), + Self::BaseEn => ("openai/whisper-base.en", "refs/pr/13"), + Self::Small => ("openai/whisper-small", "main"), + Self::SmallEn => ("openai/whisper-small.en", "refs/pr/10"), + Self::Medium => ("openai/whisper-medium", "main"), + Self::MediumEn => ("openai/whisper-medium.en", "main"), + Self::Large => ("openai/whisper-large", "refs/pr/36"), + Self::LargeV2 => ("openai/whisper-large-v2", "refs/pr/57"), + Self::LargeV3 => ("openai/whisper-large-v3", "main"), + Self::LargeV3Turbo => ("openai/whisper-large-v3-turbo", "main"), + Self::DistilMediumEn => ("distil-whisper/distil-medium.en", "main"), + Self::DistilLargeV2 => ("distil-whisper/distil-large-v2", "main"), + } + } +} + +pub struct WhisperProcessor { + pub model: WhisperModel, + pub tokenizer: Tokenizer, + pub config: Config, + pub mel_filters: Vec, + pub device: Device, +} + +impl WhisperProcessor { + pub fn new( + model: WhichWhisperModel, + device: Device, + manager: progress::ModelLoadProgressEmitterManager, + ) -> Result { + // Load the Whisper model based on the provided model type + let api = ApiBuilder::new().with_progress(false).build()?; + let (model_id, revision) = model.model_and_revision(); + let repo = api.repo(Repo::with_revision( + model_id.to_string(), + RepoType::Model, + revision.to_string(), + )); + + let config_filename = + repo.download_with_progress("config.json", manager.clone().new_processor("config.json"))?; + println!("config_filename: {:?}", config_filename.display()); + let tokenizer_filename = repo.download_with_progress( + "tokenizer.json", + manager.clone().new_processor("tokenizer.json"), + )?; + println!("tokenizer_filename: {:?}", tokenizer_filename.display()); + let model_filename = repo.download_with_progress( + "model.safetensors", + manager.clone().new_processor("model.safetensors"), + )?; + println!("model_filename: {:?}", model_filename.display()); + + let config: Config = serde_json::from_str(&std::fs::read_to_string(config_filename)?)?; + let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(anyhow::Error::msg)?; + + println!("Loading Whisper model from: {:?}", model_filename.display()); + + // SAFETY: This is safe because we are using a mmaped file and the safetensors library guarantees that the data is valid. + let var_builder = unsafe { + VarBuilder::from_mmaped_safetensors(&[model_filename], whisper_model::DTYPE, &device)? + }; + + let model = WhisperModel::Normal(whisper_model::model::Whisper::load( + &var_builder, + config.clone(), + )?); + + let mel_bytes = match config.num_mel_bins { + 80 => include_bytes!("./melfilters.bytes").as_slice(), + 128 => include_bytes!("./melfilters128.bytes").as_slice(), + num_mel_bins => anyhow::bail!("Unsupported number of mel bins: {}", num_mel_bins), + }; + + let mut mel_filters = vec![0f32; mel_bytes.len() / 4]; + ::read_f32_into(mel_bytes, &mut mel_filters); + + Ok(Self { + model, + tokenizer, + config, + mel_filters, + device, + }) + } + + pub fn transcribe( + &mut self, + audio: &[f32], + ) -> Result { + // Convert PCM to mel spectrogram + let mel = audio::pcm_to_mel(&self.config, audio, &self.mel_filters); + let mel_len = mel.len(); + let mel = Tensor::from_vec( + mel, + ( + 1, + self.config.num_mel_bins, + mel_len / self.config.num_mel_bins, + ), + &self.device, + )?; + + // Run inference + let audio_features = self.model.encoder_forward(&mel, true)?; + // Simple greedy decoding + let tokens = self.decode_greedy(&audio_features)?; + + let text = self + .tokenizer + .decode(&tokens, true) + .map_err(anyhow::Error::msg)?; + + Ok(text) + } + + fn decode_greedy( + &mut self, + audio_features: &Tensor, + ) -> Result> { + let mut tokens = vec![ + self.token_id(whisper_model::SOT_TOKEN)?, + self.token_id(whisper_model::TRANSCRIBE_TOKEN)?, + self.token_id(whisper_model::NO_TIMESTAMPS_TOKEN)?, + ]; + + let max_len = 50; // Short sequence for real-time processing + + for i in 0..max_len { + let tokens_t = Tensor::new(tokens.as_slice(), &self.device)?.unsqueeze(0)?; + let ys = self + .model + .decoder_forward(&tokens_t, audio_features, i == 0)?; + + let (_, seq_len, _) = ys.dims3()?; + let logits = self + .model + .decoder_final_linear(&ys.i((..1, seq_len - 1..))?)? + .i(0)? + .i(0)?; + + // Get most likely token + let logits_v: Vec = logits.to_vec1()?; + let next_token = logits_v + .iter() + .enumerate() + .max_by(|(_, a), (_, b)| a.total_cmp(b)) + .map(|(i, _)| u32::try_from(i).unwrap()) + .unwrap(); + + if next_token == self.token_id(whisper_model::EOT_TOKEN)? { + break; + } + + tokens.push(next_token); + } + + Ok(tokens) + } + + fn token_id( + &self, + token: &str, + ) -> Result { + self + .tokenizer + .token_to_id(token) + .ok_or_else(|| anyhow::anyhow!("Token not found: {}", token)) + } +} diff --git a/apps/stage-tamagotchi/src/pages/index.vue b/apps/stage-tamagotchi/src/pages/index.vue index ea5e592d9..5f773c76a 100644 --- a/apps/stage-tamagotchi/src/pages/index.vue +++ b/apps/stage-tamagotchi/src/pages/index.vue @@ -1,6 +1,5 @@