| 1 | use std::path::{Path, PathBuf}; |
| 2 | use std::sync::{Arc, Mutex}; |
| 3 | use whisper_rs::{FullParams, SamplingStrategy, WhisperContext, WhisperContextParameters}; |
| 4 | |
| 5 | |
| 6 | #[derive(Default)] |
| 7 | pub struct WhisperCache { |
| 8 | loaded: Mutex<Option<(PathBuf, Arc<WhisperContext>)>>, |
| 9 | } |
| 10 | |
| 11 | impl WhisperCache { |
| 12 | fn context(&self, model_path: &Path) -> Result<Arc<WhisperContext>, String> { |
| 13 | let mut guard = self.loaded.lock().unwrap(); |
| 14 | if let Some((path, ctx)) = guard.as_ref() { |
| 15 | if path == model_path { |
| 16 | return Ok(ctx.clone()); |
| 17 | } |
| 18 | } |
| 19 | let ctx = WhisperContext::new_with_params( |
| 20 | model_path |
| 21 | .to_str() |
| 22 | .ok_or("model path is not valid UTF-8".to_string())?, |
| 23 | WhisperContextParameters::default(), |
| 24 | ) |
| 25 | .map_err(|e| format!("load whisper model: {e}"))?; |
| 26 | let ctx = Arc::new(ctx); |
| 27 | *guard = Some((model_path.to_path_buf(), ctx.clone())); |
| 28 | Ok(ctx) |
| 29 | } |
| 30 | |
| 31 | pub fn transcribe( |
| 32 | &self, |
| 33 | model_path: &Path, |
| 34 | samples: &[f32], |
| 35 | language: &str, |
| 36 | ) -> Result<String, String> { |
| 37 | let ctx = self.context(model_path)?; |
| 38 | let mut state = ctx |
| 39 | .create_state() |
| 40 | .map_err(|e| format!("create whisper state: {e}"))?; |
| 41 | |
| 42 | let mut params = FullParams::new(SamplingStrategy::Greedy { best_of: 1 }); |
| 43 | let lang = if language.is_empty() { "auto" } else { language }; |
| 44 | params.set_language(Some(lang)); |
| 45 | params.set_print_special(false); |
| 46 | params.set_print_progress(false); |
| 47 | params.set_print_realtime(false); |
| 48 | params.set_print_timestamps(false); |
| 49 | params.set_suppress_blank(true); |
| 50 | let threads = std::thread::available_parallelism() |
| 51 | .map(|n| n.get()) |
| 52 | .unwrap_or(4) |
| 53 | .min(8) as i32; |
| 54 | params.set_n_threads(threads); |
| 55 | |
| 56 | state |
| 57 | .full(params, samples) |
| 58 | .map_err(|e| format!("transcription failed: {e}"))?; |
| 59 | |
| 60 | let mut text = String::new(); |
| 61 | for segment in state.as_iter() { |
| 62 | if let Ok(piece) = segment.to_str_lossy() { |
| 63 | text.push_str(&piece); |
| 64 | } |
| 65 | } |
| 66 | Ok(text.trim().to_string()) |
| 67 | } |
| 68 | } |