Skip to content

Commit ad932e8

Browse files
committed
fix: major bug when swapping STT engines durring runtime.
1 parent cde91ba commit ad932e8

8 files changed

Lines changed: 217 additions & 61 deletions

File tree

‎crates/voxctrl-config/src/lib.rs‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ pub enum ConfigError {
1515

1616
// ── Engine sub-configs ────────────────────────────────────────────────────────
1717

18-
#[derive(Debug, Clone, Serialize, Deserialize)]
18+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
1919
pub struct WhisperCppConfig {
2020
/// Directory containing GGUF model files. Empty = platform default.
2121
pub model_dir: String,
@@ -42,7 +42,7 @@ impl Default for WhisperCppConfig {
4242
}
4343
}
4444

45-
#[derive(Debug, Clone, Serialize, Deserialize)]
45+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
4646
pub struct MoonshineConfig {
4747
/// "base" or "tiny"
4848
pub model_size: String,
@@ -59,7 +59,7 @@ impl Default for MoonshineConfig {
5959
}
6060
}
6161

62-
#[derive(Debug, Clone, Serialize, Deserialize)]
62+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
6363
pub struct ParakeetConfig {
6464
pub model_size: String,
6565
pub language: String,
@@ -74,7 +74,7 @@ impl Default for ParakeetConfig {
7474
}
7575
}
7676

77-
#[derive(Debug, Clone, Serialize, Deserialize)]
77+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
7878
pub struct RemoteOpenAiConfig {
7979
/// Remote OpenAI-compatible endpoint URL, e.g. "http://localhost:8000/v1"
8080
pub endpoint: String,
@@ -120,7 +120,7 @@ impl Default for BackendChoice {
120120
}
121121
}
122122

123-
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
123+
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
124124
pub struct EngineConfig {
125125
#[serde(default)]
126126
pub backend: BackendChoice,

‎crates/voxctrl-inference/src/lib.rs‎

Lines changed: 123 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -252,6 +252,34 @@ impl InferenceEngine {
252252
self.backend.unload();
253253
}
254254

255+
/// Update engine configuration. If backend or backend model settings changed,
256+
/// re-creates the backend and returns `true` (meaning the caller should reload).
257+
pub fn update_config(&mut self, new_config: Arc<AppConfig>) -> bool {
258+
let backend_changed = self.config.engine.backend != new_config.engine.backend
259+
|| (new_config.engine.backend == BackendChoice::WhisperCpp
260+
&& self.config.engine.whisper_cpp != new_config.engine.whisper_cpp)
261+
|| (new_config.engine.backend == BackendChoice::Moonshine
262+
&& self.config.engine.moonshine != new_config.engine.moonshine)
263+
|| (new_config.engine.backend == BackendChoice::Parakeet
264+
&& self.config.engine.parakeet != new_config.engine.parakeet)
265+
|| (new_config.engine.backend == BackendChoice::RemoteOpenAi
266+
&& self.config.engine.remote_openai != new_config.engine.remote_openai);
267+
268+
self.config = new_config.clone();
269+
270+
if backend_changed {
271+
info!(
272+
"Inference backend configuration changed, switching backend to {:?}",
273+
new_config.engine.backend
274+
);
275+
self.backend.unload();
276+
self.backend = build_backend(&new_config);
277+
true
278+
} else {
279+
false
280+
}
281+
}
282+
255283
/// Transcribe and post-process. Returns the final text.
256284
pub fn process(&self, req: InferenceRequest) -> Result<InferenceOutput> {
257285
if req.audio.is_empty() {
@@ -533,6 +561,19 @@ pub fn run_worker(
533561
config: Arc<AppConfig>,
534562
rx: Receiver<InferenceRequest>,
535563
tx: Sender<InferenceOutput>,
564+
) {
565+
let (_dummy_tx, dummy_rx) = crossbeam_channel::unbounded();
566+
run_worker_with_config(config, rx, tx, dummy_rx);
567+
}
568+
569+
/// Run the inference engine on a dedicated OS thread with dynamic config reloading.
570+
/// Receives `InferenceRequest` from `rx`, sends `InferenceOutput` to `tx`,
571+
/// and updates/reloads the backend whenever `config_rx` receives a new `AppConfig`.
572+
pub fn run_worker_with_config(
573+
config: Arc<AppConfig>,
574+
rx: Receiver<InferenceRequest>,
575+
tx: Sender<InferenceOutput>,
576+
config_rx: Receiver<Arc<AppConfig>>,
536577
) {
537578
std::thread::Builder::new()
538579
.name("voxctrl-inference".into())
@@ -554,42 +595,70 @@ pub fn run_worker(
554595
}
555596
};
556597

557-
while let Ok(req) = rx.recv() {
558-
if !loaded {
559-
match engine.load() {
560-
Ok(()) => {
561-
info!("Inference engine ready (loaded on demand)");
562-
loaded = true;
563-
}
564-
Err(e) => {
565-
error!("Inference backend still not loadable: {e:#}");
566-
let _ = tx.send(InferenceOutput {
567-
text: String::new(),
568-
target_id: req.target_id,
569-
raw_text: String::new(),
570-
inference_ms: 0,
571-
language: String::new(),
572-
error: Some(format!("{e:#}")),
573-
});
574-
continue;
598+
loop {
599+
crossbeam_channel::select! {
600+
recv(rx) -> req_res => {
601+
let req = match req_res {
602+
Ok(r) => r,
603+
Err(_) => break,
604+
};
605+
606+
if !loaded {
607+
match engine.load() {
608+
Ok(()) => {
609+
info!("Inference engine ready (loaded on demand)");
610+
loaded = true;
611+
}
612+
Err(e) => {
613+
error!("Inference backend still not loadable: {e:#}");
614+
let _ = tx.send(InferenceOutput {
615+
text: String::new(),
616+
target_id: req.target_id,
617+
raw_text: String::new(),
618+
inference_ms: 0,
619+
language: String::new(),
620+
error: Some(format!("{e:#}")),
621+
});
622+
continue;
623+
}
624+
}
575625
}
576-
}
577-
}
578626

579-
match engine.process(req) {
580-
Ok(output) => {
581-
let _ = tx.send(output);
627+
match engine.process(req) {
628+
Ok(output) => {
629+
let _ = tx.send(output);
630+
}
631+
Err(e) => {
632+
error!("Inference error: {:?}", e);
633+
let _ = tx.send(InferenceOutput {
634+
text: "".to_string(),
635+
target_id: "".to_string(),
636+
raw_text: "".to_string(),
637+
inference_ms: 0,
638+
language: "".to_string(),
639+
error: Some(format!("{e:#}")),
640+
});
641+
}
642+
}
582643
}
583-
Err(e) => {
584-
error!("Inference error: {:?}", e);
585-
let _ = tx.send(InferenceOutput {
586-
text: "".to_string(),
587-
target_id: "".to_string(),
588-
raw_text: "".to_string(),
589-
inference_ms: 0,
590-
language: "".to_string(),
591-
error: Some(format!("{e:#}")),
592-
});
644+
recv(config_rx) -> new_cfg_res => {
645+
let new_cfg = match new_cfg_res {
646+
Ok(c) => c,
647+
Err(_) => break,
648+
};
649+
let needs_reload = engine.update_config(new_cfg);
650+
if needs_reload {
651+
loaded = match engine.load() {
652+
Ok(()) => {
653+
info!("Inference engine ready with new backend");
654+
true
655+
}
656+
Err(e) => {
657+
error!("Failed to load new inference backend: {e:#}");
658+
false
659+
}
660+
};
661+
}
593662
}
594663
}
595664
}
@@ -649,4 +718,25 @@ mod tests {
649718
assert_eq!(backend.name(), "remote-openai");
650719
assert!(backend.is_loaded());
651720
}
721+
722+
#[test]
723+
fn test_engine_update_config_switches_backend() {
724+
let cfg = AppConfig::default();
725+
let mut engine = InferenceEngine::new(Arc::new(cfg.clone()));
726+
assert_eq!(engine.backend.name(), "whisper-cpp");
727+
728+
let mut new_cfg = cfg.clone();
729+
new_cfg.engine.backend = BackendChoice::RemoteOpenAi;
730+
new_cfg.engine.remote_openai.endpoint = "http://localhost:5000/v1".to_string();
731+
let reloaded = engine.update_config(Arc::new(new_cfg));
732+
assert!(reloaded);
733+
assert_eq!(engine.backend.name(), "remote-openai");
734+
735+
// Non-backend config change should not trigger backend reload
736+
let mut features_cfg = engine.config.as_ref().clone();
737+
features_cfg.features.remove_fillers = !features_cfg.features.remove_fillers;
738+
let reloaded_features = engine.update_config(Arc::new(features_cfg));
739+
assert!(!reloaded_features);
740+
assert_eq!(engine.backend.name(), "remote-openai");
741+
}
652742
}

‎crates/voxctrl-inference/src/parakeet.rs‎

Lines changed: 67 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -170,31 +170,19 @@ fn load_vocab(path: &Path) -> Result<Vec<String>> {
170170
}
171171

172172
fn detokenize(tokens: &[usize], vocab: &[String]) -> String {
173-
let mut text = String::new();
173+
let mut raw = String::new();
174174
for &tok_id in tokens {
175175
if tok_id >= vocab.len() {
176176
continue;
177177
}
178178
let tok = &vocab[tok_id];
179-
if tok == "<blk>" || tok == "<unk>" || tok.is_empty() {
179+
if tok == "<blk>" || tok == "<unk>" || tok == "<pad>" || tok.starts_with("<|") || tok.is_empty() {
180180
continue;
181181
}
182-
// Handle SentencePiece whitespace prefix ( or Ġ)
183-
if let Some(stripped) = tok.strip_prefix(' ') {
184-
if !text.is_empty() {
185-
text.push(' ');
186-
}
187-
text.push_str(stripped);
188-
} else if let Some(stripped) = tok.strip_prefix('Ġ') {
189-
if !text.is_empty() {
190-
text.push(' ');
191-
}
192-
text.push_str(stripped);
193-
} else {
194-
text.push_str(tok);
195-
}
182+
raw.push_str(tok);
196183
}
197-
text
184+
let converted = raw.replace('\u{2581}', " ").replace('Ġ', " ");
185+
converted.trim().to_string()
198186
}
199187

200188
// ── Loaded State ──────────────────────────────────────────────────────────────
@@ -206,6 +194,9 @@ struct Loaded {
206194
vocab: Vec<String>,
207195
targets_is_i32: bool,
208196
decoder_enc_shape_time_first: bool,
197+
logits_idx: usize,
198+
state_1_idx: usize,
199+
state_2_idx: usize,
209200
}
210201

211202
// ── Backend ───────────────────────────────────────────────────────────────────
@@ -331,13 +322,34 @@ impl TranscriptionBackend for ParakeetBackend {
331322
})
332323
.unwrap_or(false);
333324

325+
let logits_idx = decoder
326+
.outputs()
327+
.iter()
328+
.position(|o| o.name() == "outputs")
329+
.unwrap_or(0);
330+
331+
let state_1_idx = decoder
332+
.outputs()
333+
.iter()
334+
.position(|o| o.name() == "output_states_1")
335+
.unwrap_or(2);
336+
337+
let state_2_idx = decoder
338+
.outputs()
339+
.iter()
340+
.position(|o| o.name() == "output_states_2")
341+
.unwrap_or(3);
342+
334343
*self.state.lock().unwrap() = Some(Loaded {
335344
preprocessor,
336345
encoder,
337346
decoder,
338347
vocab,
339348
targets_is_i32,
340349
decoder_enc_shape_time_first,
350+
logits_idx,
351+
state_1_idx,
352+
state_2_idx,
341353
});
342354
self.loaded = true;
343355
Ok(())
@@ -472,7 +484,7 @@ fn run_inference(state: &mut Loaded, audio: &[f32]) -> Result<String> {
472484
}
473485

474486
// ── 3. Decoder: TDT Greedy Search Loop ───────────────────────────────────
475-
let vocab_size = state.vocab.len().min(BLANK_TOKEN_ID);
487+
let vocab_size = state.vocab.len();
476488
let output_dim = vocab_size + NUM_DURATION_CLASSES;
477489
let blank_idx = BLANK_TOKEN_ID;
478490

@@ -547,7 +559,7 @@ fn run_inference(state: &mut Loaded, audio: &[f32]) -> Result<String> {
547559
];
548560

549561
let dec_out = state.decoder.run(dec_feed).context("decoder step run")?;
550-
let (_, ldata) = dec_out[0]
562+
let (_, ldata) = dec_out[state.logits_idx]
551563
.try_extract_tensor::<f32>()
552564
.context("extract decoder logits")?;
553565

@@ -567,10 +579,10 @@ fn run_inference(state: &mut Loaded, audio: &[f32]) -> Result<String> {
567579
emitted_tokens.push(best_token);
568580
current_token = best_token;
569581

570-
let (_, next_s1) = dec_out[1]
582+
let (_, next_s1) = dec_out[state.state_1_idx]
571583
.try_extract_tensor::<f32>()
572584
.context("extract next state_1")?;
573-
let (_, next_s2) = dec_out[2]
585+
let (_, next_s2) = dec_out[state.state_2_idx]
574586
.try_extract_tensor::<f32>()
575587
.context("extract next state_2")?;
576588
state_1 = next_s1.to_vec();
@@ -614,4 +626,38 @@ mod tests {
614626
let text = detokenize(&tokens, &vocab);
615627
assert_eq!(text, "Hello world!");
616628
}
629+
#[test]
630+
fn test_inspect_decoder() {
631+
let path = std::path::Path::new("/home/jrufer/.local/share/voxctrl/models/parakeet/tdt-0.6b-v3/decoder_joint-model.int8.onnx");
632+
if !path.exists() {
633+
return;
634+
}
635+
let session = ParakeetBackend::build_session(path).unwrap();
636+
let outputs = session.outputs();
637+
let s1_idx = outputs.iter().position(|o| o.name() == "output_states_1").unwrap_or(2);
638+
let s2_idx = outputs.iter().position(|o| o.name() == "output_states_2").unwrap_or(3);
639+
let logits_idx = outputs.iter().position(|o| o.name() == "outputs").unwrap_or(0);
640+
assert_eq!(logits_idx, 0);
641+
assert_eq!(s1_idx, 2);
642+
assert_eq!(s2_idx, 3);
643+
}
644+
645+
#[test]
646+
fn test_parakeet_transcribe_silence() {
647+
let _dir = model_size_dir("", "tdt-0.6b-v3");
648+
if !is_model_downloaded("tdt-0.6b-v3", "") {
649+
return;
650+
}
651+
let mut backend = ParakeetBackend::new(ParakeetConfig::default());
652+
backend.load().expect("load parakeet");
653+
let req = TranscribeRequest {
654+
audio: vec![0.0f32; 16000],
655+
language: None,
656+
word_timestamps: false,
657+
initial_prompt: None,
658+
};
659+
let res = backend.transcribe(&req).expect("transcribe silence");
660+
println!("Parakeet transcribe silence result: {:?}", res.text);
661+
assert_eq!(res.text, "");
662+
}
617663
}

‎src-tauri/src/commands.rs‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -148,6 +148,9 @@ pub async fn save_config(
148148
guard.save().map_err(|e| e.to_string())?;
149149
info!("Config saved");
150150

151+
// Hot-reload inference engine configuration
152+
let _ = state.inference_config_tx.send(Arc::new(new_config.clone()));
153+
151154
let (overlay_position, overlay_monitor) = (
152155
guard.data.ui.overlay_position.clone(),
153156
guard.data.ui.overlay_monitor.clone(),

0 commit comments

Comments
 (0)