voice_bird_cli/transcription/
refinement_engine.rs1use std::path::PathBuf;
2use std::time::Duration;
3
4use tokio::sync::{broadcast, mpsc, oneshot};
5use whisper_rs::{FullParams, SamplingStrategy, WhisperContext, WhisperContextParameters};
6
7use super::{EngineEvent, EngineHandle, Segment, Token};
8
9pub struct RefinementEngine {
14 pub model_path: PathBuf,
15 pub language: Option<String>,
16 pub window_ms: u32,
17 pub beam_size: u8,
18}
19
20impl RefinementEngine {
21 pub fn start(self) -> anyhow::Result<EngineHandle> {
22 let (pcm_tx, mut pcm_rx) = mpsc::channel::<Vec<f32>>(32);
23 let (events_tx, events_rx) = broadcast::channel::<EngineEvent>(256);
24 let (shutdown_tx, mut shutdown_rx) = oneshot::channel::<()>();
25
26 let model_path = self.model_path;
27 let language = self.language;
28 let window_ms = self.window_ms.max(5_000) as u64;
29 let beam_size = self.beam_size.max(1) as i32;
30 const OVERLAP_MS: u64 = 1_000;
33
34 std::thread::spawn(move || {
35 let ctx = match WhisperContext::new_with_params(
36 model_path.to_string_lossy().as_ref(),
37 WhisperContextParameters::default(),
38 ) {
39 Ok(c) => c,
40 Err(e) => {
41 let _ =
42 events_tx.send(EngineEvent::Error(format!("refinement load model: {e}")));
43 return;
44 }
45 };
46 let _ = events_tx.send(EngineEvent::ModelLoaded {
47 name: format!(
48 "refine:{}",
49 model_path.file_name().unwrap_or_default().to_string_lossy()
50 ),
51 });
52
53 let mut state = match ctx.create_state() {
54 Ok(s) => s,
55 Err(e) => {
56 let _ =
57 events_tx.send(EngineEvent::Error(format!("refinement create state: {e}")));
58 return;
59 }
60 };
61
62 let mut buffer: Vec<f32> = Vec::new();
63 let mut abs_offset_ms: u64 = 0;
65
66 loop {
67 if shutdown_rx.try_recv().is_ok() {
68 break;
69 }
70
71 let mut end_of_stream = false;
72 match pcm_rx.blocking_recv() {
73 Some(chunk) => buffer.extend_from_slice(&chunk),
74 None => end_of_stream = true,
75 }
76
77 let window_samples = ((window_ms * 16_000) / 1000) as usize;
78 let overlap_samples = ((OVERLAP_MS * 16_000) / 1000) as usize;
79
80 while buffer.len() >= window_samples {
83 let take = (window_samples + overlap_samples).min(buffer.len());
84 let slice: Vec<f32> = buffer[..take].to_vec();
85
86 run_pass(
87 &mut state,
88 &slice,
89 abs_offset_ms,
90 abs_offset_ms + window_ms,
93 &language,
94 beam_size,
95 &events_tx,
96 );
97
98 buffer.drain(..window_samples);
99 abs_offset_ms += window_ms;
100 }
101
102 if end_of_stream {
103 const MIN_FLUSH_MS: u64 = 2_000;
107 let tail_ms = (buffer.len() as u64 * 1000) / 16_000;
108 if tail_ms >= MIN_FLUSH_MS {
109 let slice: Vec<f32> = buffer.clone();
110 run_pass(
111 &mut state,
112 &slice,
113 abs_offset_ms,
114 abs_offset_ms + tail_ms,
115 &language,
116 beam_size,
117 &events_tx,
118 );
119 }
120 break;
121 }
122 }
123 });
124
125 Ok(EngineHandle {
126 pcm_tx,
127 events_rx,
128 shutdown: shutdown_tx,
129 })
130 }
131}
132
133#[allow(clippy::too_many_arguments)]
134fn run_pass(
135 state: &mut whisper_rs::WhisperState,
136 input: &[f32],
137 abs_t_start_ms: u64,
138 abs_t_end_ms: u64,
139 language: &Option<String>,
140 beam_size: i32,
141 events_tx: &broadcast::Sender<EngineEvent>,
142) {
143 let mut params = FullParams::new(SamplingStrategy::BeamSearch {
144 beam_size,
145 patience: -1.0,
146 });
147 params.set_no_context(true);
148 params.set_print_progress(false);
149 params.set_print_realtime(false);
150 params.set_print_special(false);
151 params.set_print_timestamps(false);
152 params.set_token_timestamps(true);
153 if let Some(ref lang) = language {
154 params.set_language(Some(lang.as_str()));
155 }
156
157 const WHISPER_MIN_SAMPLES: usize = 16_000 + 1_600;
159 let padded;
160 let inference_input: &[f32] = if input.len() < WHISPER_MIN_SAMPLES {
161 padded = {
162 let mut v = Vec::with_capacity(WHISPER_MIN_SAMPLES);
163 v.extend_from_slice(input);
164 v.resize(WHISPER_MIN_SAMPLES, 0.0);
165 v
166 };
167 &padded
168 } else {
169 input
170 };
171
172 let inf_start = std::time::Instant::now();
173 if let Err(e) = state.full(params, inference_input) {
174 let _ = events_tx.send(EngineEvent::Error(format!("refinement full: {e}")));
175 return;
176 }
177 let inf_ms = inf_start.elapsed().as_millis() as u64;
178 let buf_ms_now = (inference_input.len() as u64 * 1000) / 16_000;
179 log::info!(
180 "refinement inference: buf={}ms took={}ms (rt_ratio={:.2})",
181 buf_ms_now,
182 inf_ms,
183 inf_ms as f64 / buf_ms_now.max(1) as f64
184 );
185
186 let n_segments = state.full_n_segments().unwrap_or(0);
187 let mut tokens: Vec<Token> = Vec::new();
188 for i in 0..n_segments {
189 let n_tokens = state.full_n_tokens(i).unwrap_or(0);
190 for t in 0..n_tokens {
191 let txt = match state.full_get_token_text(i, t) {
192 Ok(s) => s,
193 Err(_) => continue,
194 };
195 if txt.starts_with("[_") {
196 continue;
197 }
198 let data = match state.full_get_token_data(i, t) {
199 Ok(d) => d,
200 Err(_) => continue,
201 };
202 let t0 = data.t0.max(0) as u64 * 10;
203 let t1 = data.t1.max(0) as u64 * 10;
204 tokens.push(Token {
205 text: txt,
206 t_start_ms: abs_t_start_ms + t0,
207 t_end_ms: abs_t_start_ms + t1,
208 });
209 }
210 }
211
212 if tokens.is_empty() {
213 return;
214 }
215
216 let text = tokens
217 .iter()
218 .map(|t| t.text.trim())
219 .filter(|s| !s.is_empty())
220 .collect::<Vec<_>>()
221 .join(" ");
222 if text.is_empty() {
223 return;
224 }
225
226 let seg = Segment {
227 t_start: Duration::from_millis(abs_t_start_ms),
228 t_end: Duration::from_millis(abs_t_end_ms),
229 text,
230 tokens,
231 };
232 let _ = events_tx.send(EngineEvent::Committed(seg));
233}