From d05fce516f38fb750b3f09575036f115967cce10 Mon Sep 17 00:00:00 2001 From: RainbowBird Date: Sat, 28 Jun 2025 00:50:01 +0800 Subject: [PATCH] feat: load whisper (#234) Co-authored-by: Neko --- apps/stage-tamagotchi/src-tauri/src/lib.rs | 30 ++- .../src-tauri/src/whisper/melfilters.bytes | Bin 0 -> 64320 bytes .../src-tauri/src/whisper/melfilters128.bytes | Bin 0 -> 102912 bytes .../src-tauri/src/whisper/mod.rs | 207 ++++++++++++++++++ 4 files changed, 234 insertions(+), 3 deletions(-) create mode 100644 apps/stage-tamagotchi/src-tauri/src/whisper/melfilters.bytes create mode 100644 apps/stage-tamagotchi/src-tauri/src/whisper/melfilters128.bytes create mode 100644 apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs diff --git a/apps/stage-tamagotchi/src-tauri/src/lib.rs b/apps/stage-tamagotchi/src-tauri/src/lib.rs index 966792de3..9fb81a774 100644 --- a/apps/stage-tamagotchi/src-tauri/src/lib.rs +++ b/apps/stage-tamagotchi/src-tauri/src/lib.rs @@ -1,13 +1,13 @@ use std::{sync::atomic::Ordering, time::Duration}; use tauri::{ - menu::{Menu, MenuItem}, - tray::TrayIconBuilder, Emitter, Manager, RunEvent, WebviewUrl, WebviewWindowBuilder, + menu::{Menu, MenuItem}, + tray::TrayIconBuilder, }; use tauri_plugin_prevent_default::Flags; use tokio::time::sleep; @@ -15,12 +15,13 @@ use tokio::time::sleep; mod app_click_through; mod app_windows; mod commands; +mod whisper; #[cfg(target_os = "macos")] use app_click_through::native_macos::{get_mouse_location, get_window_frame}; #[cfg(target_os = "windows")] use app_click_through::native_windows::{get_mouse_location, get_window_frame}; -use app_click_through::state::{set_click_through_enabled, WindowClickThroughState}; +use app_click_through::state::{WindowClickThroughState, set_click_through_enabled}; #[tauri::command] async fn start_monitor(window: tauri::Window) -> Result<(), String> { @@ -95,6 +96,26 @@ async fn stop_click_through(window: tauri::Window) -> Result<(), String> { Ok(()) } +fn load_whisper_model() -> Result { + // Determine device to use + let device = if candle_core::utils::cuda_is_available() { + candle_core::Device::new_cuda(0)? + } else if candle_core::utils::metal_is_available() { + candle_core::Device::new_metal(0)? + } else { + candle_core::Device::Cpu + }; + + println!("Using device: {device:?}"); + + let whisper_model = whisper::WhichWhisperModel::Tiny; + + println!("Loading whisper model: {:?}", whisper_model); + + let model = whisper::WhisperProcessor::new(whisper_model, device.clone())?; + Ok(model) +} + #[cfg_attr(mobile, tauri::mobile_entry_point)] #[allow(clippy::missing_panics_doc)] pub fn run() { @@ -140,6 +161,9 @@ 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>)?; diff --git a/apps/stage-tamagotchi/src-tauri/src/whisper/melfilters.bytes b/apps/stage-tamagotchi/src-tauri/src/whisper/melfilters.bytes new file mode 100644 index 0000000000000000000000000000000000000000..0874829e2088c94e3a4f001725e7040fae57fa0c GIT binary patch literal 64320 zcmeI*c~p;S9|rKmB$Oqxy%HkL2rZU=_jQ|^;_V${O=THd)?TtC+k1veYBZT?K@2rx zsYsbO6kb^(lq@qE(}~G4Lue#3rWtGWHuJaNUH{*ga~$V9-{+t2=UmshpXb-_#KgqJ zBz50KS@#eBP=EqP1VUDtB5}$w_QZ$}FH_(zfsmCESl2R>Jz*Rupb*%Xa|Zv(+s2+K zj5wk|-E)_mw>XA9VH7B!5QuAUCy#h8Wlt1F98q9jj+fk)I*vVI6eyq&IPWt{hMKvu zCki8uC{Xu#mkbEBVow+a3Md5D6mF28ZMqA^EgVrmdx5z2DRT6LUF?bWNg#kq;Op7h z($US6Jy9WXroBMj&l%;tH=#m%2m(+*A&}en2fo8Ad^n=OUjo}wud*kM0tFNTOM|1t z=YD6{6NM2+6mWDNhETI?_JmQOfI{HP4=Gr5WEXp)Fye>;=56jFAaDbF!YEKcA+XcF ztsK`Wf;~|faYTWWXI<;Qvoe!CVH7B!5cqM1uMB$X$(|^TIHJHalVx?^?drjvFbWh< z2*i9BBVBSEu_p>6jwrCTYf{~JAoKYhNJfDI3V|238S?Ylk?e`Wh$9O8(jof^V zziWj_pnytX=;1MPPqYVrexgF+OnZTwu^nX6q!9K*`y>!RC9w8UHmWK%$)-_md`D41 zai+CENY(|h>Rgm;*{hm8(K-oSzCM9mr@`{v!=Zfkst<~Dtp%><6(BNmoy@rx%QHpm zByf3^z^Ip>inSOkA2he+nWBQ?OnU*B$K?o_79!os<9McMp9TUL6F7hKs7P~_a@pu! zvQbH~!PtGgP66EoQbr%Zk_EvsW^)P86y3wX{gnd!cQut)dw(bwG>KqWlu{h(EZ`r| z3g6zam+!2%lV#Ve*%h6`z}=Ms6E?4gf3BOXIMiG^CuQ;fdr)d|sIx#^^;fX-cbA?% zFYt?3B)g(>7`VGq;Ml8lB(ImU`_NzETbsnLD784$S>TaJ3W{r-WtVsVK+&a#?267| z;O#xuIL;F?yeN5h_ZsY z&mE*kcbBURvhdN(nS2gYigBo`z^JXhV$#@joNM|I`6~D>zOx?0vqjf1aPN8rlJdu) zpl22SoMtIkElkG3ip4xz>Q&> zAEb;bKtNa<*)8KRYBSogFAb98O}Yt$+if<~dd)=P&(|^bnTZ^q{1y5;+~c#DZb9I_ z4HEF2zC?t4vl{aPZo_8pV{AB`fX;ze*_Q_C@h05_j?UgHir0i;ro(j%thx{L@^~zo zU&+4c76k6wAc5k7SYi7z2=o6{guJ{8Jh9pe?@LyEe`$~&Z_-U5|NIBy=7kS&F5v(= z%=#U}lG0J0Is`Q%9C*g)76$HHufUX-E{506o$%hajqp$U5sr;cB5Ps@EN7IkFZJs2 z1|0=z%#Mq;p9a7$@*q0CEJfq&FHu(Nf>ZWm*%=)J!JU-?V|JGqd_&vevE4f4n4iV7 z%ejag5RM*$THtP{{ruf8N<9vB6sVaRDNLWc!1JdBL=L@xs`KgS(QGzKOr8jbPu$rV z9YevLl>*=I*(j91Gai*E`O*%_rChdK)EwAm-B3I`$b z)OxhFIgaHX+i@r;0IMc6LHjl~e4o)V6x>-U(61mt#4YQArsk{QemMif3scd2atOlS zwZVg`FW4ES9)~/utqUOq5I&oVza%-w|i&@2o{T!-P~-SM&cZ4nXun$KrC27@~r z71-lC(Qxa0j@UJ@H%iTx!YpV9TxM)V`)h$Hs_csAZ}y667iadys2~SA3Pf+}Bg$Ks zi~K$YME4Gd)6og z(#K#(wKc>ihsB?^rukIIf+EGB9UQggJuza$T3-mW6?1%&7O)lUk40+ zbXRPjx==*f@8PqWF+^V1PM}~{uAz9u7V+e_S`j^b0KQBPgzt)2%wDn-&pLYIv%FRq zcqvz8yEyX9(JmcaSS1kZ)L+~R`=2r2f&GaC!Y6-_ zh)C7JgT3)n7SF0PsHa1y!wG*PNa}duT?UKTU-zFfvPZ?&+m?Soi$q-K$ tJrJ4RR+wDsgqCwhBlOlt#4mF~ibY$@^R5)}%l3uGe zLq~@uhZc<*G=`)YgBpjhhiS7$jfOE>dp_#DpT6IJUs%ue_;CH!|6135=T!!SVWr=x zA;Q346rg}ffwH<6oDq{cEKuNY0l!nloD+tD0uq5ENt-z%5+9Z*P%qG{d7HkJpa2C- z3k1~Wp?~Bpn0|v53Q*ukk}Y=$V?Y5Nft(dPcsJ28Vp(kgzZ(Ud5w#=0?i(R6yV}4r zS0kMGin;;;wLan(_c%>m0d`ISeSzgCx^b82Gw}_@1^P!$&WN0d zRSFmeq+aEmFb5Qn2pr9e;fzRpSfW6^K+XzB-rpGo3P=Qk<|M+&!oYnZF=B}Vp53gu zOBe$R=m=y_*~49;W5lxB0&hM)&lyoW0_?sK0{@w8&NEjdocM~m0)suey!k9w3;tg$ zbtAye8zJy2urqf_Bb@k(+5#eQ0q-ViM}Xby2wdy^kTarV#4-ime2)mu31dJ3i9l}o zD$a<+hb0Qs3wS*3!23I+KmmzB+Nr5%_d_f06NwQ^6gZk^$z8%2P(Vi@%RiI5M8}9_ zwFL&*W^+c=jsUxFgutvT&3NW&gcDy;S0HGPov0WzhM(n9Hv;Uu5dsm{ow-XI;lx+e z7I3my$-9Z#5n%T^0zZzZ;Ed=Pu}pz})2lcqi~$8C0t>3bI3p4tmMBm!kUgan@9&HP z1tbE)%pLK&YZvYli4jW_FkjY$yM!^IfR4cT_7}KIbc|S5TVQCHWdh@V;?KoWI|l4t zN5C`8o@Xu{BbF(UzH>LvT#Nw)Bm)0E93#Ha`tdW55+jx<5FXu>XD-Ho0y+ZiGtThL zrDMdh+5*|PJjJw=?wk{~W5Dip1ajM*;Ed=Pu}py@US_5nRCJ@P(UJ(oYr0I9(a~>A~9l#0#zAP;8dBxIbjqiAQ3o{n1%5jQ#mIR zBbF!-TvP_fNn1E4i~?}6w(2_p2-Y2HyJPApL%Jt;&S3U}q%i&3C}M4;p& zD{-w)n0OT$&9j%pi6ylK-uW~fbCzxpt7f0)*-PyxuzMW=%Y*$z`?+2`d+9i_thPYz zAEj`~jS>ZqVz^V(jsm-v2owY)YwMeg6q~!W=T4D0u|xrjy${jr!eS8;ypB7CaiD++ zfoZYdY10e*#qtnm?i3S7d`y9Qfnhc0urqj>7@eQXox(^^z=XiD)&)j~t{p|%yP7DO z-;}dr!ikS5V08Wx58^%$iLPcme=!aekO*uwyN!$1Q$@&!Iov4{CzjL}2#)vF)()_J z^RrPj;n0WQ&tB~`uzQ(6h?x=M&mLlPTw7u3^pKzZl3}r?v_S05jWBNa6*m&z5l%zg zI4er0fz8VVlB|xx=}EAto_UXV6&Vz3$_tFxRg7LwMvL^&OPm$u(?9@I0+G$ndJncP zM`F!jQIb)}Suq90=gJG1&)Q=w-BpIWfdL}azKFA;d?E;7O5l!Lvhmw(#kd_DAoeY| zg#E8t@osMli_eu7xbYI+zOC}$ch_5_?5TtZ%Hgaiod`BB6IgDWf`ItW!t2cMIG*po zbC?W^HKhfHxGqCVPaE;i5It;GQG7jh3x<2g)Mj4!Av z5L@SpB8v(<@7h+BTW6xVdn$K}s$pQ?Qh}{m0h-;pBsk8q7GEEFjBAU=a<@puSX5cS zzQ7%oZ+=ew$T4=J&FO4-1>fXuQ8^Io-K4;Q#r9f6={oeDYY=|(Y7x*c5*y6!7)|~K z7AT;wz`#9qT2^Tmnx?cANA~@IPWSurET(WC*t%RG^|$@cl-fvy4=J5n`?1(@L|r2qIqE3a)B=fp49$Py%u4X zcQB#Z&j`0#hv0taIWKZHRuvVXj@#vNJob#e+9@w^AAauq$ z?OgN>c(@m0q4#6-)OMiq^S1o@n4FDOMFm!DcGgOgeDTrjOc-2V;95mG#`X4t&F)^@ zF^Xn_ZJQSOD{6#~hnGEE6BFQLU4fUg3eauRBAh5Y$9XZ$#)`rM(_R*8(O)e<=Ae9B z8Fm9po+rTD%NdrAh5Ss6!nt7U4GXM%?X1N=b;IfGZSa|W3%O4YWBQh}@hWQCuL%E=GH{$_FX065*eB84i~F@qXxJ3>|u#J4W$r zu>FPwhTlJ_W!D8^hanjQnwG+H)D9fLFof?t%b97Ij87F7sD4@Fy*{Q~+mtdIPmzKt zb1qo_Py3BQ^8%+Fdd&NED-qN7R^1~0h-$^bR3%mU;mv5dax9>{_i8tFJE*1b0hCH u3g?5Zzb)|1j5_0>->+(6#h&Q5?Ms+@?tx|VP1tBV;mz+h^?&~VuK6Ejlbd+} literal 0 HcmV?d00001 diff --git a/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs b/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs new file mode 100644 index 000000000..a60d1a39b --- /dev/null +++ b/apps/stage-tamagotchi/src-tauri/src/whisper/mod.rs @@ -0,0 +1,207 @@ +// 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}; +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, +} + +impl WhisperProcessor { + pub fn new( + model: WhichWhisperModel, + device: Device, + ) -> Result { + // Load the Whisper model based on the provided model type + let api = Api::new()?; + 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.get("config.json")?; + let tokenizer_filename = repo.get("tokenizer.json")?; + 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)) + } +}