Skip to main content

voice_bird_cli/transcription/
refinement_engine.rs

1use 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
9/// Background refinement engine. Runs on non-overlapping windows with
10/// beam search for higher-quality transcripts than the real-time
11/// streaming engine. Emits `Committed` segments only — no `Tentative`,
12/// since refinement output replaces streaming output wholesale.
13pub 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        // Carry a small overlap across chunks so sentences spanning a
31        // window boundary aren't split mid-word.
32        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            // Absolute session time of buffer[0].
64            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                // Process as many full windows as the buffer holds. Each
81                // pass emits one refined segment.
82                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                        // The refined segment's absolute range is the
91                        // *committed* window, not the overlap tail.
92                        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                    // Flush whatever's left if it's long enough to be worth
104                    // decoding (< 1 s will be rejected by whisper.cpp, and
105                    // very short tails produce hallucinations).
106                    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    // whisper.cpp rejects < 1 s inputs.
158    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}