From cdd4c1a845528670e9f17b22de1f53d3cf0ba811 Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Tue, 7 Jul 2026 09:39:10 -0700 Subject: [PATCH 1/5] feat(voxtral): add realtime stream buffer --- crates/voice-voxtral/src/lib.rs | 4 + crates/voice-voxtral/src/realtime_stream.rs | 464 ++++++++++++++++++++ 2 files changed, 468 insertions(+) create mode 100644 crates/voice-voxtral/src/realtime_stream.rs diff --git a/crates/voice-voxtral/src/lib.rs b/crates/voice-voxtral/src/lib.rs index 35b8ae5..efb788c 100644 --- a/crates/voice-voxtral/src/lib.rs +++ b/crates/voice-voxtral/src/lib.rs @@ -17,6 +17,7 @@ mod prompt; mod realtime; mod realtime_audio; mod realtime_inference; +mod realtime_stream; mod streaming; mod text; mod tokenizer; @@ -79,6 +80,9 @@ pub use realtime_inference::{ VoxtralRealtimeTokenEmbeddings, VoxtralRealtimeTranscriber, VoxtralRealtimeTranscription, VoxtralRealtimeTranscriptionOptions, }; +pub use realtime_stream::{ + VoxtralRealtimeStreamBuffer, VoxtralRealtimeStreamConfig, VoxtralRealtimeStreamWindow, +}; pub use streaming::{ plan_codec_chunk, VoxtralCodecChunk, VoxtralStreamingConfig, DEFAULT_CODEC_CHUNK_FRAMES, DEFAULT_CODEC_CHUNK_FRAMES_AT_BEGIN, DEFAULT_CODEC_LEFT_CONTEXT_FRAMES, diff --git a/crates/voice-voxtral/src/realtime_stream.rs b/crates/voice-voxtral/src/realtime_stream.rs new file mode 100644 index 0000000..dfc04c9 --- /dev/null +++ b/crates/voice-voxtral/src/realtime_stream.rs @@ -0,0 +1,464 @@ +use std::collections::VecDeque; + +use crate::{ + build_realtime_streaming_prompt_with_left_pad, realtime_num_delay_tokens, + realtime_raw_audio_length_per_token, Result, VoxtralError, VoxtralRealtimeConfig, + VoxtralTokenizerMetadata, REALTIME_DEFAULT_OFFLINE_BUFFER_TOKENS, +}; + +/// Model-native scheduling parameters for Voxtral Realtime STT. +/// +/// This mirrors the upstream realtime buffer contract: a first prompt window +/// consumes the streaming prefix, then every generated text token is fed back +/// while the audio side advances by one raw-audio token. +#[derive(Debug, Clone, PartialEq)] +pub struct VoxtralRealtimeStreamConfig { + pub sample_rate: u32, + pub raw_audio_length_per_token: usize, + pub look_ahead_samples: usize, + pub look_back_samples: usize, + pub left_pad_tokens: usize, + pub delay_tokens: usize, + pub right_pad_tokens: usize, + pub prompt_token_ids: Vec, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct VoxtralRealtimeStreamWindow { + pub sequence: usize, + pub input_token_ids: Vec, + pub audio_samples: Vec, + pub frame_start_sample: usize, + pub frame_end_sample: usize, + pub stride_start_sample: usize, + pub stride_end_sample: usize, + pub input_token_start: usize, + pub input_token_end: usize, + pub is_initial: bool, +} + +#[derive(Debug, Clone)] +pub struct VoxtralRealtimeStreamBuffer { + config: VoxtralRealtimeStreamConfig, + samples: Vec, + token_queue: VecDeque, + stride_start_sample: usize, + stride_end_sample: usize, + input_token_position: usize, + sequence: usize, + finished: bool, +} + +impl VoxtralRealtimeStreamConfig { + pub fn from_metadata( + config: &VoxtralRealtimeConfig, + tokenizer: &VoxtralTokenizerMetadata, + ) -> Result { + let delay_ms = tokenizer.audio.transcription_delay_ms.ok_or_else(|| { + VoxtralError::InvalidTokenizer("missing realtime transcription_delay_ms".into()) + })?; + let delay_tokens = realtime_num_delay_tokens(config, delay_ms)?; + Self::from_metadata_with_delay_tokens(config, tokenizer, delay_tokens) + } + + pub fn from_metadata_with_delay_tokens( + config: &VoxtralRealtimeConfig, + tokenizer: &VoxtralTokenizerMetadata, + delay_tokens: usize, + ) -> Result { + let raw_audio_length_per_token = realtime_raw_audio_length_per_token(config)?; + let sample_rate = config.sample_rate(); + if tokenizer.audio.sampling_rate != sample_rate { + return Err(VoxtralError::InvalidTokenizer(format!( + "tokenizer sampling_rate={} but realtime config sample_rate={sample_rate}", + tokenizer.audio.sampling_rate + ))); + } + if (tokenizer.audio.frame_rate - config.frame_rate()).abs() > f64::EPSILON { + return Err(VoxtralError::InvalidTokenizer(format!( + "tokenizer frame_rate={} but realtime config frame_rate={}", + tokenizer.audio.frame_rate, + config.frame_rate() + ))); + } + + let left_pad_tokens = tokenizer.audio.streaming_n_left_pad_tokens.ok_or_else(|| { + VoxtralError::InvalidTokenizer("missing realtime streaming_n_left_pad_tokens".into()) + })?; + let look_ahead_samples = ms_to_samples( + tokenizer.audio.streaming_look_ahead_ms.ok_or_else(|| { + VoxtralError::InvalidTokenizer("missing realtime streaming_look_ahead_ms".into()) + })?, + sample_rate, + )?; + let look_back_samples = ms_to_samples( + tokenizer.audio.streaming_look_back_ms.ok_or_else(|| { + VoxtralError::InvalidTokenizer("missing realtime streaming_look_back_ms".into()) + })?, + sample_rate, + )?; + let right_pad_tokens = delay_tokens + .checked_add(1) + .and_then(|tokens| tokens.checked_add(REALTIME_DEFAULT_OFFLINE_BUFFER_TOKENS)) + .ok_or_else(|| VoxtralError::InvalidConfig("right pad token overflow".into()))?; + let prompt = build_realtime_streaming_prompt_with_left_pad(left_pad_tokens, delay_tokens); + + Ok(Self { + sample_rate, + raw_audio_length_per_token, + look_ahead_samples, + look_back_samples, + left_pad_tokens, + delay_tokens, + right_pad_tokens, + prompt_token_ids: prompt.input_ids, + }) + } + + pub fn left_pad_samples(&self) -> usize { + self.left_pad_tokens * self.raw_audio_length_per_token + } + + pub fn right_pad_samples(&self) -> usize { + self.right_pad_tokens * self.raw_audio_length_per_token + } + + pub fn initial_stride_end_sample(&self) -> usize { + self.prompt_token_ids.len() * self.raw_audio_length_per_token + } +} + +impl VoxtralRealtimeStreamBuffer { + pub fn new(config: VoxtralRealtimeStreamConfig) -> Result { + if config.raw_audio_length_per_token == 0 { + return Err(VoxtralError::InvalidConfig( + "raw audio length per token must be greater than zero".into(), + )); + } + if config.prompt_token_ids.is_empty() { + return Err(VoxtralError::InvalidConfig( + "realtime stream prompt must not be empty".into(), + )); + } + + let mut samples = Vec::new(); + samples.resize(config.left_pad_samples(), 0.0); + let token_queue = config.prompt_token_ids.iter().copied().collect(); + let stride_end_sample = config.initial_stride_end_sample(); + + Ok(Self { + config, + samples, + token_queue, + stride_start_sample: 0, + stride_end_sample, + input_token_position: 0, + sequence: 0, + finished: false, + }) + } + + pub fn config(&self) -> &VoxtralRealtimeStreamConfig { + &self.config + } + + pub fn buffered_samples(&self) -> usize { + self.samples.len() + } + + pub fn queued_tokens(&self) -> usize { + self.token_queue.len() + } + + pub fn push_audio_16khz(&mut self, samples: &[f32]) -> Result<()> { + if self.finished { + return Err(VoxtralError::InvalidConfig( + "cannot push audio after realtime stream finish".into(), + )); + } + self.samples.extend_from_slice(samples); + Ok(()) + } + + pub fn push_generated_token(&mut self, token: usize) { + self.token_queue.push_back(token); + } + + pub fn finish(&mut self) { + if self.finished { + return; + } + let align_pad_samples = (self.config.raw_audio_length_per_token + - (self.samples.len() % self.config.raw_audio_length_per_token)) + % self.config.raw_audio_length_per_token; + self.samples + .extend(std::iter::repeat_n(0.0, align_pad_samples)); + self.samples + .extend(std::iter::repeat_n(0.0, self.config.right_pad_samples())); + self.finished = true; + } + + pub fn next_window(&mut self) -> Result> { + let stride_samples = self + .stride_end_sample + .checked_sub(self.stride_start_sample) + .ok_or_else(|| VoxtralError::InvalidConfig("invalid realtime stream stride".into()))?; + if !stride_samples.is_multiple_of(self.config.raw_audio_length_per_token) { + return Err(VoxtralError::InvalidConfig(format!( + "stream stride {stride_samples} is not divisible by raw audio token size {}", + self.config.raw_audio_length_per_token + ))); + } + let token_count = stride_samples / self.config.raw_audio_length_per_token; + if self.token_queue.len() < token_count { + return Ok(None); + } + + let frame_start = self + .stride_start_sample + .saturating_sub(self.config.look_back_samples); + let frame_end = self + .stride_end_sample + .checked_add(self.config.look_ahead_samples) + .ok_or_else(|| VoxtralError::InvalidConfig("stream frame end overflow".into()))?; + if self.samples.len() < frame_end { + return Ok(None); + } + + let input_token_ids = self.token_queue.drain(..token_count).collect::>(); + let input_token_start = self.input_token_position; + let input_token_end = input_token_start + token_count; + let window = VoxtralRealtimeStreamWindow { + sequence: self.sequence, + input_token_ids, + audio_samples: self.samples[frame_start..frame_end].to_vec(), + frame_start_sample: frame_start, + frame_end_sample: frame_end, + stride_start_sample: self.stride_start_sample, + stride_end_sample: self.stride_end_sample, + input_token_start, + input_token_end, + is_initial: self.sequence == 0, + }; + + self.sequence += 1; + self.input_token_position = input_token_end; + self.stride_start_sample = self.stride_end_sample; + self.stride_end_sample = self + .stride_end_sample + .checked_add(self.config.raw_audio_length_per_token) + .ok_or_else(|| VoxtralError::InvalidConfig("stream stride overflow".into()))?; + + Ok(Some(window)) + } +} + +fn ms_to_samples(ms: f64, sample_rate: u32) -> Result { + if !ms.is_finite() || ms < 0.0 { + return Err(VoxtralError::InvalidConfig(format!( + "streaming millisecond value must be finite and non-negative, got {ms}" + ))); + } + let samples = ms * sample_rate as f64 / 1000.0; + let rounded = samples.round(); + if (samples - rounded).abs() > 1e-6 { + return Err(VoxtralError::InvalidConfig(format!( + "{ms}ms at {sample_rate}Hz does not map to an integral sample count" + ))); + } + Ok(rounded as usize) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + REALTIME_BOS_TOKEN_ID, REALTIME_SAMPLE_RATE, REALTIME_STREAMING_PAD_TOKEN_ID, + REALTIME_TRANSCRIPTION_FORMAT, + }; + + fn tiny_realtime_config() -> VoxtralRealtimeConfig { + VoxtralRealtimeConfig { + dim: 8, + n_layers: 1, + head_dim: 4, + hidden_dim: 16, + n_heads: 2, + n_kv_heads: 1, + use_biases: false, + causal: true, + rope_theta: 1_000_000.0, + norm_eps: 1e-5, + vocab_size: 64, + model_parallel: 1, + tied_embeddings: true, + sliding_window: 32, + model_max_length: 128, + multimodal: crate::VoxtralRealtimeMultimodalConfig { + whisper_model_args: crate::VoxtralRealtimeWhisperModelConfig { + encoder_args: crate::VoxtralRealtimeAudioEncoderConfig { + audio_encoding_args: crate::VoxtralRealtimeAudioEncodingConfig { + sampling_rate: REALTIME_SAMPLE_RATE, + frame_rate: 12.5, + num_mel_bins: 4, + hop_length: 160, + window_size: 400, + chunk_length_s: None, + global_log_mel_max: 1.5, + transcription_format: REALTIME_TRANSCRIPTION_FORMAT.to_string(), + }, + dim: 8, + n_layers: 1, + head_dim: 4, + hidden_dim: 16, + n_heads: 2, + vocab_size: 64, + n_kv_heads: 1, + use_biases: true, + use_cache: false, + rope_theta: 1_000_000.0, + causal: true, + norm_eps: 1e-5, + pos_embed: "rope".to_string(), + max_source_positions: None, + ffn_type: "swiglu".to_string(), + norm_type: "rms_norm".to_string(), + sliding_window: 16, + }, + downsample_args: crate::VoxtralRealtimeDownsampleConfig { + downsample_factor: 4, + }, + }, + }, + ada_rms_norm_t_cond: true, + ada_rms_norm_t_cond_dim: Some(32), + } + } + + fn tokenizer_json() -> String { + serde_json::json!({ + "config": { + "pattern": ".", + "num_vocab_tokens": 0, + "default_vocab_size": 64, + "default_num_special_tokens": 35, + "version": "v1" + }, + "vocab": [], + "special_tokens": (0..35).map(|rank| { + let token = match rank { + 1 => "", + 2 => "", + 24 => "[AUDIO]", + 25 => "[BEGIN_AUDIO]", + 26 => "[OUTPUT_AUDIO]", + 32 => "[STREAMING_PAD]", + 33 => "[STREAMING_WORD]", + 34 => "[REPEAT_AUDIO_TEXT]", + _ => "", + }; + serde_json::json!({ + "rank": rank, + "token_str": token, + "is_control": true + }) + }).collect::>(), + "audio": { + "sampling_rate": REALTIME_SAMPLE_RATE, + "frame_rate": 12.5, + "audio_encoding_config": { + "num_mel_bins": 4, + "hop_length": 160, + "window_size": 400 + }, + "transcription_delay_ms": 480, + "streaming_look_ahead_ms": 2.5, + "streaming_look_back_ms": 52.5, + "streaming_n_left_pad_tokens": 32, + "transcription_format": REALTIME_TRANSCRIPTION_FORMAT + } + }) + .to_string() + } + + fn tokenizer() -> VoxtralTokenizerMetadata { + VoxtralTokenizerMetadata::from_json_str(&tokenizer_json()).unwrap() + } + + #[test] + fn builds_stream_config_from_realtime_metadata() { + let config = tiny_realtime_config(); + let stream = VoxtralRealtimeStreamConfig::from_metadata(&config, &tokenizer()).unwrap(); + + assert_eq!(stream.sample_rate, REALTIME_SAMPLE_RATE); + assert_eq!(stream.raw_audio_length_per_token, 1280); + assert_eq!(stream.look_ahead_samples, 40); + assert_eq!(stream.look_back_samples, 840); + assert_eq!(stream.left_pad_tokens, 32); + assert_eq!(stream.delay_tokens, 6); + assert_eq!(stream.right_pad_tokens, 17); + assert_eq!(stream.left_pad_samples(), 40_960); + assert_eq!(stream.right_pad_samples(), 21_760); + assert_eq!(stream.initial_stride_end_sample(), 49_920); + assert_eq!(stream.prompt_token_ids.len(), 39); + assert_eq!(stream.prompt_token_ids[0], REALTIME_BOS_TOKEN_ID); + assert!(stream.prompt_token_ids[1..] + .iter() + .all(|token| *token == REALTIME_STREAMING_PAD_TOKEN_ID)); + } + + #[test] + fn emits_initial_window_then_waits_for_audio_and_token_feedback() { + let config = tiny_realtime_config(); + let stream_config = + VoxtralRealtimeStreamConfig::from_metadata(&config, &tokenizer()).unwrap(); + let mut buffer = VoxtralRealtimeStreamBuffer::new(stream_config).unwrap(); + + assert_eq!(buffer.buffered_samples(), 40_960); + assert_eq!(buffer.queued_tokens(), 39); + assert!(buffer.next_window().unwrap().is_none()); + + buffer.push_audio_16khz(&vec![0.5; 9_000]).unwrap(); + let initial = buffer.next_window().unwrap().unwrap(); + assert!(initial.is_initial); + assert_eq!(initial.sequence, 0); + assert_eq!(initial.input_token_ids.len(), 39); + assert_eq!(initial.frame_start_sample, 0); + assert_eq!(initial.frame_end_sample, 49_960); + assert_eq!(initial.stride_start_sample, 0); + assert_eq!(initial.stride_end_sample, 49_920); + assert_eq!(initial.audio_samples.len(), 49_960); + + assert!(buffer.next_window().unwrap().is_none()); + buffer.push_generated_token(41); + assert!(buffer.next_window().unwrap().is_none()); + + buffer.push_audio_16khz(&vec![0.25; 1_280]).unwrap(); + let next = buffer.next_window().unwrap().unwrap(); + assert!(!next.is_initial); + assert_eq!(next.sequence, 1); + assert_eq!(next.input_token_ids, vec![41]); + assert_eq!(next.frame_start_sample, 49_080); + assert_eq!(next.frame_end_sample, 51_240); + assert_eq!(next.stride_start_sample, 49_920); + assert_eq!(next.stride_end_sample, 51_200); + assert_eq!(next.audio_samples.len(), 2_160); + } + + #[test] + fn finish_adds_alignment_and_right_padding_once() { + let config = tiny_realtime_config(); + let stream_config = + VoxtralRealtimeStreamConfig::from_metadata(&config, &tokenizer()).unwrap(); + let mut buffer = VoxtralRealtimeStreamBuffer::new(stream_config).unwrap(); + buffer.push_audio_16khz(&vec![1.0; 1_281]).unwrap(); + + buffer.finish(); + let once = buffer.buffered_samples(); + buffer.finish(); + + assert_eq!(buffer.buffered_samples(), once); + assert_eq!(once, 40_960 + 1_281 + 1_279 + 21_760); + assert!(buffer.push_audio_16khz(&[1.0]).is_err()); + } +} From 41500c8340a58b475c433dd3818c185a16cdc38e Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Tue, 7 Jul 2026 09:44:23 -0700 Subject: [PATCH 2/5] feat(voxtral): add realtime stream session --- crates/voice-voxtral/src/lib.rs | 2 + crates/voice-voxtral/src/realtime_session.rs | 375 +++++++++++++++++++ 2 files changed, 377 insertions(+) create mode 100644 crates/voice-voxtral/src/realtime_session.rs diff --git a/crates/voice-voxtral/src/lib.rs b/crates/voice-voxtral/src/lib.rs index efb788c..2dd84a8 100644 --- a/crates/voice-voxtral/src/lib.rs +++ b/crates/voice-voxtral/src/lib.rs @@ -17,6 +17,7 @@ mod prompt; mod realtime; mod realtime_audio; mod realtime_inference; +mod realtime_session; mod realtime_stream; mod streaming; mod text; @@ -80,6 +81,7 @@ pub use realtime_inference::{ VoxtralRealtimeTokenEmbeddings, VoxtralRealtimeTranscriber, VoxtralRealtimeTranscription, VoxtralRealtimeTranscriptionOptions, }; +pub use realtime_session::{VoxtralRealtimeStreamSession, VoxtralRealtimeStreamStep}; pub use realtime_stream::{ VoxtralRealtimeStreamBuffer, VoxtralRealtimeStreamConfig, VoxtralRealtimeStreamWindow, }; diff --git a/crates/voice-voxtral/src/realtime_session.rs b/crates/voice-voxtral/src/realtime_session.rs new file mode 100644 index 0000000..7d9dc22 --- /dev/null +++ b/crates/voice-voxtral/src/realtime_session.rs @@ -0,0 +1,375 @@ +use candle_core::Tensor; + +use crate::{ + realtime_log_mel_spectrogram, Result, VoxtralError, VoxtralRealtimeStreamBuffer, + VoxtralRealtimeStreamConfig, VoxtralRealtimeStreamWindow, VoxtralRealtimeTranscriber, + VoxtralRealtimeTranscriptionOptions, REALTIME_EOS_TOKEN_ID, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct VoxtralRealtimeStreamStep { + pub sequence: usize, + pub token: usize, + pub text: String, + pub reached_eos: bool, + pub output_token_count: usize, + pub frame_start_sample: usize, + pub frame_end_sample: usize, + pub input_token_start: usize, + pub input_token_end: usize, +} + +/// Stateful Voxtral Realtime decoding session. +/// +/// This implements the model-native token/audio feedback loop. It intentionally +/// does not cache decoder KV state yet; each step recomputes the current text +/// history while preserving the externally visible streaming contract. +pub struct VoxtralRealtimeStreamSession<'a> { + transcriber: &'a VoxtralRealtimeTranscriber, + buffer: VoxtralRealtimeStreamBuffer, + options: VoxtralRealtimeTranscriptionOptions, + input_token_ids: Vec, + output_tokens: Vec, + audio_embeddings: Option, + done: bool, +} + +impl VoxtralRealtimeTranscriber { + pub fn stream_session( + &self, + config: VoxtralRealtimeStreamConfig, + options: VoxtralRealtimeTranscriptionOptions, + ) -> Result> { + VoxtralRealtimeStreamSession::new(self, config, options) + } +} + +impl<'a> VoxtralRealtimeStreamSession<'a> { + pub fn new( + transcriber: &'a VoxtralRealtimeTranscriber, + config: VoxtralRealtimeStreamConfig, + options: VoxtralRealtimeTranscriptionOptions, + ) -> Result { + Ok(Self { + transcriber, + buffer: VoxtralRealtimeStreamBuffer::new(config)?, + options, + input_token_ids: Vec::new(), + output_tokens: Vec::new(), + audio_embeddings: None, + done: false, + }) + } + + pub fn push_audio_16khz(&mut self, samples: &[f32]) -> Result<()> { + self.buffer.push_audio_16khz(samples) + } + + pub fn push_generated_token_for_test(&mut self, token: usize) { + self.buffer.push_generated_token(token); + } + + pub fn finish(&mut self) { + self.buffer.finish(); + } + + pub fn output_tokens(&self) -> &[usize] { + &self.output_tokens + } + + pub fn text(&self) -> String { + decode_text( + &self.transcriber.token_decoder, + self.output_tokens.iter().copied(), + ) + } + + pub fn next_step(&mut self) -> Result> { + if self.done || self.output_tokens.len() >= self.options.max_new_tokens { + return Ok(None); + } + let Some(window) = self.buffer.next_window()? else { + return Ok(None); + }; + + let window_embeddings = self.transcriber.encode_stream_window_embeddings(&window)?; + self.audio_embeddings = Some(match &self.audio_embeddings { + Some(existing) => { + Tensor::cat(&[existing, &window_embeddings], 1).map_err(candle_err)? + } + None => window_embeddings, + }); + self.input_token_ids + .extend(window.input_token_ids.iter().copied()); + + let audio_embeddings = self + .audio_embeddings + .as_ref() + .expect("audio embeddings are set above"); + let generation = self + .transcriber + .text_decoder + .greedy_decode_audio_embeddings_with_prompt( + &self.transcriber.token_embeddings, + audio_embeddings, + &self.input_token_ids, + self.options.delay_tokens, + 1, + ) + .map_err(candle_err)?; + let Some(token) = generation.generated_tokens.first().copied() else { + return Ok(None); + }; + + let reached_eos = token == REALTIME_EOS_TOKEN_ID; + self.output_tokens.push(token); + if reached_eos { + self.done = true; + } else { + self.buffer.push_generated_token(token); + } + let text = self.text(); + + Ok(Some(VoxtralRealtimeStreamStep { + sequence: window.sequence, + token, + text, + reached_eos, + output_token_count: self.output_tokens.len(), + frame_start_sample: window.frame_start_sample, + frame_end_sample: window.frame_end_sample, + input_token_start: window.input_token_start, + input_token_end: window.input_token_end, + })) + } +} + +impl VoxtralRealtimeTranscriber { + pub fn encode_stream_window_embeddings( + &self, + window: &VoxtralRealtimeStreamWindow, + ) -> Result { + let mel = realtime_log_mel_spectrogram(&self.config, &window.audio_samples)?; + let input_features = Tensor::from_vec( + mel.to_channel_major(), + (1, mel.mel_bins, mel.frames), + self.token_embeddings.tok_embeddings.embeddings().device(), + ) + .map_err(candle_err)? + .to_dtype(self.token_embeddings.tok_embeddings.embeddings().dtype()) + .map_err(candle_err)?; + let audio_start_pos = window + .input_token_start + .checked_mul(self.config.downsample_factor()) + .ok_or_else(|| VoxtralError::InvalidConfig("stream audio position overflow".into()))?; + let embeddings = self + .audio_modules + .forward(&input_features, audio_start_pos) + .map_err(candle_err)?; + let actual_tokens = embeddings.dim(1).map_err(candle_err)?; + let expected_tokens = window.input_token_ids.len(); + if actual_tokens < expected_tokens { + return Err(VoxtralError::Candle(format!( + "stream window produced {actual_tokens} audio embeddings for {expected_tokens} input tokens" + ))); + } + if actual_tokens == expected_tokens { + return Ok(embeddings); + } + embeddings + .narrow(1, actual_tokens - expected_tokens, expected_tokens) + .map_err(candle_err) + } +} + +fn candle_err(err: candle_core::Error) -> VoxtralError { + VoxtralError::Candle(err.to_string()) +} + +fn decode_text( + decoder: &crate::VoxtralTekkenDecoder, + tokens: impl IntoIterator, +) -> String { + let text_tokens = tokens + .into_iter() + .filter(|token| *token != REALTIME_EOS_TOKEN_ID) + .collect::>(); + decoder.decode(&text_tokens).trim().to_string() +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + VoxtralRealtimeAudioModules, VoxtralRealtimeConfig, VoxtralRealtimeDownsampleConfig, + VoxtralRealtimeInferenceModules, VoxtralRealtimeMultimodalConfig, + VoxtralRealtimeWhisperModelConfig, VoxtralTokenizerMetadata, REALTIME_SAMPLE_RATE, + REALTIME_TRANSCRIPTION_FORMAT, + }; + use candle_core::{DType, Device}; + use candle_nn::VarBuilder; + + fn tiny_realtime_config() -> VoxtralRealtimeConfig { + VoxtralRealtimeConfig { + dim: 8, + n_layers: 1, + head_dim: 4, + hidden_dim: 16, + n_heads: 2, + n_kv_heads: 1, + use_biases: false, + causal: true, + rope_theta: 1_000_000.0, + norm_eps: 1e-5, + vocab_size: 64, + model_parallel: 1, + tied_embeddings: true, + sliding_window: 32, + model_max_length: 128, + multimodal: VoxtralRealtimeMultimodalConfig { + whisper_model_args: VoxtralRealtimeWhisperModelConfig { + encoder_args: crate::VoxtralRealtimeAudioEncoderConfig { + audio_encoding_args: crate::VoxtralRealtimeAudioEncodingConfig { + sampling_rate: REALTIME_SAMPLE_RATE, + frame_rate: 12.5, + num_mel_bins: 4, + hop_length: 160, + window_size: 400, + chunk_length_s: None, + global_log_mel_max: 1.5, + transcription_format: REALTIME_TRANSCRIPTION_FORMAT.to_string(), + }, + dim: 8, + n_layers: 1, + head_dim: 4, + hidden_dim: 16, + n_heads: 2, + vocab_size: 64, + n_kv_heads: 1, + use_biases: true, + use_cache: false, + rope_theta: 1_000_000.0, + causal: true, + norm_eps: 1e-5, + pos_embed: "rope".to_string(), + max_source_positions: None, + ffn_type: "swiglu".to_string(), + norm_type: "rms_norm".to_string(), + sliding_window: 16, + }, + downsample_args: VoxtralRealtimeDownsampleConfig { + downsample_factor: 4, + }, + }, + }, + ada_rms_norm_t_cond: true, + ada_rms_norm_t_cond_dim: Some(32), + } + } + + fn tokenizer_json() -> String { + serde_json::json!({ + "config": { + "pattern": ".", + "num_vocab_tokens": 0, + "default_vocab_size": 64, + "default_num_special_tokens": 35, + "version": "v1" + }, + "vocab": [], + "special_tokens": (0..35).map(|rank| { + let token = match rank { + 1 => "", + 2 => "", + 24 => "[AUDIO]", + 25 => "[BEGIN_AUDIO]", + 26 => "[OUTPUT_AUDIO]", + 32 => "[STREAMING_PAD]", + 33 => "[STREAMING_WORD]", + 34 => "[REPEAT_AUDIO_TEXT]", + _ => "", + }; + serde_json::json!({ + "rank": rank, + "token_str": token, + "is_control": true + }) + }).collect::>(), + "audio": { + "sampling_rate": REALTIME_SAMPLE_RATE, + "frame_rate": 12.5, + "audio_encoding_config": { + "num_mel_bins": 4, + "hop_length": 160, + "window_size": 400 + }, + "transcription_delay_ms": 480, + "streaming_look_ahead_ms": 2.5, + "streaming_look_back_ms": 52.5, + "streaming_n_left_pad_tokens": 32, + "transcription_format": REALTIME_TRANSCRIPTION_FORMAT + } + }) + .to_string() + } + + fn transcriber() -> VoxtralRealtimeTranscriber { + let config = tiny_realtime_config(); + let modules = VoxtralRealtimeInferenceModules::load( + &config, + VarBuilder::zeros(DType::F32, &Device::Cpu), + ) + .unwrap(); + let audio_modules = + VoxtralRealtimeAudioModules::load(&config, VarBuilder::zeros(DType::F32, &Device::Cpu)) + .unwrap(); + let decoder = crate::VoxtralRealtimeTextDecoder::load( + &config, + VarBuilder::zeros(DType::F32, &Device::Cpu), + ) + .unwrap(); + let tokenizer = VoxtralTokenizerMetadata::from_json_str(&tokenizer_json()).unwrap(); + VoxtralRealtimeTranscriber::new( + config, + modules.token_embeddings, + audio_modules, + decoder, + tokenizer.decoder().unwrap(), + ) + } + + #[test] + fn streaming_session_waits_for_audio_then_emits_token_steps() { + let transcriber = transcriber(); + let tokenizer = VoxtralTokenizerMetadata::from_json_str(&tokenizer_json()).unwrap(); + let stream_config = + VoxtralRealtimeStreamConfig::from_metadata(&transcriber.config, &tokenizer).unwrap(); + let mut session = transcriber + .stream_session( + stream_config, + VoxtralRealtimeTranscriptionOptions { + delay_tokens: 6, + max_new_tokens: 2, + }, + ) + .unwrap(); + + assert!(session.next_step().unwrap().is_none()); + session.push_audio_16khz(&vec![0.0; 9_000]).unwrap(); + let first = session.next_step().unwrap().unwrap(); + assert_eq!(first.sequence, 0); + assert_eq!(first.input_token_start, 0); + assert_eq!(first.input_token_end, 39); + assert_eq!(first.output_token_count, 1); + + assert!(session.next_step().unwrap().is_none()); + session.push_audio_16khz(&vec![0.0; 1_280]).unwrap(); + let second = session.next_step().unwrap().unwrap(); + assert_eq!(second.sequence, 1); + assert_eq!(second.input_token_start, 39); + assert_eq!(second.input_token_end, 40); + assert_eq!(second.output_token_count, 2); + assert_eq!(session.output_tokens().len(), 2); + assert!(session.next_step().unwrap().is_none()); + } +} From 1855a1ae5b1f00dc38743c16fd75257cd0155e0e Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Tue, 7 Jul 2026 09:47:42 -0700 Subject: [PATCH 3/5] feat(stt): expose native streaming sessions --- crates/voice-stt/src/lib.rs | 144 ++++++++++++++++++++++++++++++++++++ 1 file changed, 144 insertions(+) diff --git a/crates/voice-stt/src/lib.rs b/crates/voice-stt/src/lib.rs index b4e75db..0004104 100644 --- a/crates/voice-stt/src/lib.rs +++ b/crates/voice-stt/src/lib.rs @@ -78,6 +78,35 @@ pub enum SttModel { pub struct VoxtralRealtimeSttModel { transcriber: voice_voxtral::VoxtralRealtimeTranscriber, options: voice_voxtral::VoxtralRealtimeTranscriptionOptions, + stream_config: voice_voxtral::VoxtralRealtimeStreamConfig, +} + +/// A single native streaming STT decode step. +#[derive(Debug, Clone)] +pub struct StreamTranscribeStep { + /// Full decoded transcript so far. + pub text: String, + /// Token generated at this streaming step. + pub token: u32, + /// Token IDs generated so far, excluding backend prompt tokens. + pub tokens: Vec, + /// Whether the backend generated its end-of-sequence token. + pub reached_eos: bool, + /// Model input sample rate for this streaming session. + pub sample_rate: u32, + /// Number of 16 kHz audio samples pushed into the model stream so far. + pub pushed_samples: usize, +} + +/// Backend-neutral native streaming STT session. +pub enum SttStreamSession<'a> { + Voxtral(VoxtralRealtimeSttStreamSession<'a>), +} + +/// Native Voxtral Realtime STT streaming session. +pub struct VoxtralRealtimeSttStreamSession<'a> { + session: voice_voxtral::VoxtralRealtimeStreamSession<'a>, + pushed_samples: usize, } impl SttBackend { @@ -123,6 +152,19 @@ impl SttModel { model.set_max_new_tokens(max_new_tokens); } } + + pub fn supports_native_streaming(&self) -> bool { + matches!(self, Self::Voxtral(_)) + } + + pub fn stream_session(&self) -> Result> { + match self { + Self::Whisper(_) => Err(SttError::Model( + "native streaming STT is not supported for Whisper".into(), + )), + Self::Voxtral(model) => model.stream_session().map(SttStreamSession::Voxtral), + } + } } impl VoxtralRealtimeSttModel { @@ -155,6 +197,98 @@ impl VoxtralRealtimeSttModel { sample_rate: voice_voxtral::REALTIME_SAMPLE_RATE, }) } + + pub fn stream_session(&self) -> Result> { + let session = self + .transcriber + .stream_session(self.stream_config.clone(), self.options) + .map_err(|e| SttError::Model(e.to_string()))?; + Ok(VoxtralRealtimeSttStreamSession { + session, + pushed_samples: 0, + }) + } +} + +impl<'a> SttStreamSession<'a> { + pub fn push_audio(&mut self, samples: &[f32], sample_rate: u32) -> Result<()> { + match self { + Self::Voxtral(session) => session.push_audio(samples, sample_rate), + } + } + + pub fn finish(&mut self) { + match self { + Self::Voxtral(session) => session.finish(), + } + } + + pub fn next_step(&mut self) -> Result> { + match self { + Self::Voxtral(session) => session.next_step(), + } + } + + pub fn drain_ready(&mut self) -> Result> { + let mut steps = Vec::new(); + while let Some(step) = self.next_step()? { + let reached_eos = step.reached_eos; + steps.push(step); + if reached_eos { + break; + } + } + Ok(steps) + } +} + +impl VoxtralRealtimeSttStreamSession<'_> { + pub fn push_audio(&mut self, samples: &[f32], sample_rate: u32) -> Result<()> { + let samples = if sample_rate != voice_voxtral::REALTIME_SAMPLE_RATE { + resample_linear(samples, sample_rate, voice_voxtral::REALTIME_SAMPLE_RATE) + } else { + samples.to_vec() + }; + self.pushed_samples = self.pushed_samples.saturating_add(samples.len()); + self.session + .push_audio_16khz(&samples) + .map_err(|e| SttError::Model(e.to_string())) + } + + pub fn finish(&mut self) { + self.session.finish(); + } + + pub fn next_step(&mut self) -> Result> { + let Some(step) = self + .session + .next_step() + .map_err(|e| SttError::Model(e.to_string()))? + else { + return Ok(None); + }; + let token = u32::try_from(step.token) + .map_err(|_| SttError::Model(format!("Voxtral token id {} exceeds u32", step.token)))?; + let tokens = self + .session + .output_tokens() + .iter() + .copied() + .map(|token| { + u32::try_from(token) + .map_err(|_| SttError::Model(format!("Voxtral token id {token} exceeds u32"))) + }) + .collect::>>()?; + + Ok(Some(StreamTranscribeStep { + text: step.text, + token, + tokens, + reached_eos: step.reached_eos, + sample_rate: voice_voxtral::REALTIME_SAMPLE_RATE, + pushed_samples: self.pushed_samples, + })) + } } /// Return the default inference device for STT. @@ -321,6 +455,15 @@ fn load_voxtral_realtime_model_on_device( let delay_tokens = model .default_delay_tokens() .map_err(|e| SttError::Model(e.to_string()))?; + let stream_config = + voice_voxtral::VoxtralRealtimeStreamConfig::from_metadata_with_delay_tokens( + model.config(), + model + .tokenizer() + .ok_or_else(|| SttError::Model("missing realtime tekken tokenizer".into()))?, + delay_tokens, + ) + .map_err(|e| SttError::Model(e.to_string()))?; let dtype = default_voxtral_dtype(&device); let transcriber = model .load_transcriber(dtype, &device) @@ -331,6 +474,7 @@ fn load_voxtral_realtime_model_on_device( delay_tokens, max_new_tokens: usize::MAX, }, + stream_config, }) } From 69e07d528c12c857eae447b946682b7bde8f25dd Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Tue, 7 Jul 2026 09:51:50 -0700 Subject: [PATCH 4/5] feat(cli): stream native Voxtral STT --- crates/voice-cli/src/cli.rs | 265 +++++++++++++++++++++++++++++---- crates/voice-cli/src/listen.rs | 132 ++++++++++++++++ 2 files changed, 371 insertions(+), 26 deletions(-) diff --git a/crates/voice-cli/src/cli.rs b/crates/voice-cli/src/cli.rs index e1c3a44..7101920 100644 --- a/crates/voice-cli/src/cli.rs +++ b/crates/voice-cli/src/cli.rs @@ -1110,6 +1110,10 @@ struct ListenArgs { #[arg(long)] continuous: bool, + /// Stream native STT tokens while recording. Defaults to Voxtral when no STT backend is selected. + #[arg(long)] + stream: bool, + /// STT backend to use locally or via STT_BACKEND. Defaults to whisper. #[arg(long = "stt-backend", value_enum)] stt_backend: Option, @@ -1770,12 +1774,21 @@ fn main() { match args.command { Some(Command::Listen(listen_args)) => { - let stt_options = stt_load_options( + let mut stt_options = stt_load_options( listen_args.stt_backend, listen_args.stt_model, listen_args.stt_max_new_tokens, ); - if listen_args.continuous { + if listen_args.stream && listen_args.continuous { + eprintln!("voice listen: --stream and --continuous cannot be combined yet"); + std::process::exit(1); + } + if listen_args.stream && !stt_selection_is_explicit(&stt_options) { + stt_options.backend = Some(voice_stt::SttBackend::Voxtral); + } + if listen_args.stream { + listen::listen_and_stream_with_options(stt_options); + } else if listen_args.continuous { if stt_selection_is_explicit(&stt_options) { listen::listen_continuous_with_options(stt_options); } else { @@ -4545,8 +4558,76 @@ fn run_realtime_live_turn( let mut partial_state = RealtimePartialSttState::default(); let partial_interval = args.stt_partials.then_some(args.partial_interval_ms); - let (samples, sample_rate) = warm_mic - .record_vad_with_progress( + let mut native_stream_transcription: Option = None; + let mut native_stream_elapsed_ms = 0u64; + let (samples, sample_rate) = if args.stt_partials && stt_model.supports_native_streaming() { + let mut stt_stream = stt_model + .stream_session() + .map_err(|e| format!("start native STT stream: {e}"))?; + let mut last_streamed_sample = 0usize; + let recording = warm_mic + .record_vad_with_progress( + args.duration.saturating_mul(1_000), + args.silence_timeout_ms, + args.vad_threshold, + args.noise_multiplier, + args.calibration_ms, + partial_interval, + |snapshot| { + let stt_started = Instant::now(); + let latest = maybe_emit_realtime_native_stream_transcription( + &mut stt_stream, + &mut last_streamed_sample, + args, + event_names, + turn_index, + &mut partial_state, + snapshot, + false, + )?; + native_stream_elapsed_ms = native_stream_elapsed_ms + .saturating_add(stt_started.elapsed().as_millis() as u64); + if let Some(latest) = latest { + native_stream_transcription = Some(latest); + } + Ok(()) + }, + ) + .map_err(|e| format!("record microphone: {e}"))?; + + let (samples, sample_rate) = recording; + if samples.len() > last_streamed_sample { + let stt_started = Instant::now(); + stt_stream + .push_audio(&samples[last_streamed_sample..], sample_rate) + .map_err(|e| format!("stream final microphone audio: {e}"))?; + native_stream_elapsed_ms = native_stream_elapsed_ms + .saturating_add(stt_started.elapsed().as_millis() as u64); + } + let stt_started = Instant::now(); + stt_stream.finish(); + let latest = emit_realtime_native_stream_steps( + &mut stt_stream, + args, + event_names, + turn_index, + &mut partial_state, + NativeStreamPartialContext { + audio_ms: samples.len().saturating_mul(1_000) as u64 / sample_rate.max(1) as u64, + capture_elapsed_ms: None, + peak: None, + threshold: None, + emit_allowed: true, + }, + )?; + native_stream_elapsed_ms = + native_stream_elapsed_ms.saturating_add(stt_started.elapsed().as_millis() as u64); + if let Some(latest) = latest { + native_stream_transcription = Some(latest); + } + (samples, sample_rate) + } else { + warm_mic.record_vad_with_progress( args.duration.saturating_mul(1_000), args.silence_timeout_ms, args.vad_threshold, @@ -4564,7 +4645,8 @@ fn run_realtime_live_turn( ) }, ) - .map_err(|e| format!("record microphone: {e}"))?; + .map_err(|e| format!("record microphone: {e}"))? + }; let stt_audio_ms = samples.len().saturating_mul(1_000) as u64 / sample_rate.max(1) as u64; let (speech_detected, peak, rms) = detect_realtime_speech(&samples, args.vad_threshold); emit_realtime_event( @@ -4630,27 +4712,33 @@ fn run_realtime_live_turn( } let stt_started = Instant::now(); - let Some(transcription) = listen::transcribe_samples(stt_model, &samples, sample_rate) else { - emit_realtime_event( - args.json, - event_names, - realtime_event( - "conversation.item.input_audio_transcription.failed", - serde_json::json!({ - "turn_index": turn_index, - "message": "transcription failed", - "audio_ms": stt_audio_ms, - }), - ), - ); - return Ok(RealtimeLiveTurnResult { - heard_speech: false, - stt_audio_ms, - stt_elapsed_ms: stt_started.elapsed().as_millis() as u64, - ..Default::default() - }); - }; - let stt_elapsed_ms = stt_started.elapsed().as_millis() as u64; + let (transcription, stt_elapsed_ms) = + if let Some(transcription) = native_stream_transcription { + (transcription, native_stream_elapsed_ms) + } else { + let Some(transcription) = listen::transcribe_samples(stt_model, &samples, sample_rate) + else { + emit_realtime_event( + args.json, + event_names, + realtime_event( + "conversation.item.input_audio_transcription.failed", + serde_json::json!({ + "turn_index": turn_index, + "message": "transcription failed", + "audio_ms": stt_audio_ms, + }), + ), + ); + return Ok(RealtimeLiveTurnResult { + heard_speech: false, + stt_audio_ms, + stt_elapsed_ms: stt_started.elapsed().as_millis() as u64, + ..Default::default() + }); + }; + (transcription, stt_started.elapsed().as_millis() as u64) + }; emit_realtime_event( args.json, event_names, @@ -4760,6 +4848,121 @@ fn emit_realtime_live_speech_started( ); } +#[derive(Debug, Clone, Copy)] +struct NativeStreamPartialContext { + audio_ms: u64, + capture_elapsed_ms: Option, + peak: Option, + threshold: Option, + emit_allowed: bool, +} + +#[allow(clippy::too_many_arguments)] +fn maybe_emit_realtime_native_stream_transcription( + stt_stream: &mut voice_stt::SttStreamSession<'_>, + last_streamed_sample: &mut usize, + args: &RealtimeArgs, + event_names: &mut Vec, + turn_index: u64, + state: &mut RealtimePartialSttState, + snapshot: listen::VadProgressSnapshot, + force_emit: bool, +) -> Result, String> { + if snapshot.samples.len() <= *last_streamed_sample { + return Ok(None); + } + let delta = &snapshot.samples[*last_streamed_sample..]; + *last_streamed_sample = snapshot.samples.len(); + stt_stream + .push_audio(delta, snapshot.sample_rate) + .map_err(|e| e.to_string())?; + + emit_realtime_live_speech_started( + args, + event_names, + turn_index, + snapshot.threshold, + snapshot.current_peak, + None, + snapshot.audio_ms, + state, + ); + + emit_realtime_native_stream_steps( + stt_stream, + args, + event_names, + turn_index, + state, + NativeStreamPartialContext { + audio_ms: snapshot.audio_ms, + capture_elapsed_ms: Some(snapshot.elapsed_ms), + peak: Some(snapshot.current_peak), + threshold: Some(snapshot.threshold), + emit_allowed: force_emit || snapshot.audio_ms >= args.partial_min_audio_ms, + }, + ) +} + +fn emit_realtime_native_stream_steps( + stt_stream: &mut voice_stt::SttStreamSession<'_>, + args: &RealtimeArgs, + event_names: &mut Vec, + turn_index: u64, + state: &mut RealtimePartialSttState, + context: NativeStreamPartialContext, +) -> Result, String> { + let mut latest = None; + for step in stt_stream.drain_ready().map_err(|e| e.to_string())? { + if step.text.trim().is_empty() { + if step.reached_eos { + break; + } + continue; + } + latest = Some(voice_stt::TranscribeResult { + text: step.text.clone(), + tokens: step.tokens.clone(), + sample_rate: step.sample_rate, + }); + if context.emit_allowed { + let Some((sequence, replaces_sequence, text)) = state.observe_text(step.text) else { + if step.reached_eos { + break; + } + continue; + }; + emit_realtime_event( + args.json, + event_names, + realtime_event( + "conversation.item.input_audio_transcription.partial", + serde_json::json!({ + "turn_index": turn_index, + "sequence": sequence, + "replaces_sequence": replaces_sequence, + "text": text, + "stable": false, + "native_streaming": true, + "tokens": step.tokens.len(), + "sample_rate": step.sample_rate, + "audio_ms": context.audio_ms, + "capture_elapsed_ms": context.capture_elapsed_ms, + "elapsed_ms": 0, + "peak": context.peak, + "threshold": context.threshold, + }), + ), + ); + } + if step.reached_eos { + break; + } + } + + Ok(latest) +} + fn maybe_emit_realtime_partial_transcription( stt_model: &mut voice_stt::SttModel, args: &RealtimeArgs, @@ -7469,6 +7672,16 @@ mod tests { assert_eq!(transcribe.stt_max_new_tokens, Some(32)); } + #[test] + fn parses_listen_stream_flag() { + let listen = Args::parse_from(["voice", "listen", "--stream"]); + let Some(Command::Listen(listen)) = listen.command else { + panic!("expected listen command"); + }; + assert!(listen.stream); + assert!(!listen.continuous); + } + #[test] fn parses_realtime_smoke_goal_shape() { let args = Args::parse_from([ diff --git a/crates/voice-cli/src/listen.rs b/crates/voice-cli/src/listen.rs index 2ea78c1..77726b2 100644 --- a/crates/voice-cli/src/listen.rs +++ b/crates/voice-cli/src/listen.rs @@ -714,6 +714,16 @@ pub fn warmup_permissions() -> Result<(), String> { /// /// Returns mono f32 samples at the device's native sample rate, plus the rate. pub fn record_until_interrupt() -> Result<(Vec, u32), String> { + record_until_interrupt_with_progress(None, |_| Ok(())) +} + +pub fn record_until_interrupt_with_progress( + progress_interval_ms: Option, + mut on_progress: F, +) -> Result<(Vec, u32), String> +where + F: FnMut(VadProgressSnapshot) -> Result<(), String>, +{ // Ding before mic so Bluetooth users hear the "ready" signal play_ding(); @@ -743,6 +753,12 @@ pub fn record_until_interrupt() -> Result<(Vec, u32), String> { // Wait for Enter key or Ctrl+C let enter_pressed = Arc::new(std::sync::atomic::AtomicBool::new(false)); let enter_clone = Arc::clone(&enter_pressed); + let start_time = Instant::now(); + let progress_interval = progress_interval_ms + .filter(|ms| *ms > 0) + .map(Duration::from_millis); + let mut last_progress_at = start_time; + let mut progress_error = None; let stdin_thread = std::thread::spawn(move || { let mut line = String::new(); @@ -754,6 +770,24 @@ pub fn record_until_interrupt() -> Result<(Vec, u32), String> { if INTERRUPTED.load(Ordering::Relaxed) || enter_pressed.load(Ordering::Relaxed) { break; } + if progress_interval.is_some_and(|interval| last_progress_at.elapsed() >= interval) { + let samples = buffer.lock().unwrap().clone(); + let audio_ms = samples.len().saturating_mul(1_000) as u64 / sample_rate.max(1) as u64; + let peak_bits = peak.load(Ordering::Relaxed); + let current_peak = f32::from_bits(peak_bits); + if let Err(error) = on_progress(VadProgressSnapshot { + samples, + sample_rate, + audio_ms, + elapsed_ms: start_time.elapsed().as_millis() as u64, + current_peak, + threshold: 0.0, + }) { + progress_error = Some(error); + break; + } + last_progress_at = Instant::now(); + } std::thread::sleep(std::time::Duration::from_millis(50)); } @@ -769,6 +803,10 @@ pub fn record_until_interrupt() -> Result<(Vec, u32), String> { log_recording_stats(&samples, sample_rate); maybe_save_recording(&samples, sample_rate); + if let Some(error) = progress_error { + return Err(error); + } + Ok((samples, sample_rate)) } @@ -1561,6 +1599,100 @@ pub fn listen_and_transcribe_with_options(options: SttLoadOptions) { } } +pub fn listen_and_stream_with_options(options: SttLoadOptions) { + let model = load_stt_with_options(options); + if !model.supports_native_streaming() { + eprintln!("Native streaming STT is only available for the Voxtral backend."); + std::process::exit(1); + } + + let mut stream = match model.stream_session() { + Ok(stream) => stream, + Err(e) => { + eprintln!("Failed to start STT stream: {e}"); + std::process::exit(1); + } + }; + let mut last_sample_count = 0usize; + let mut last_text = String::new(); + let mut last_tokens = 0usize; + + let record_result = record_until_interrupt_with_progress(Some(80), |snapshot| { + if snapshot.samples.len() <= last_sample_count { + return Ok(()); + } + let delta = &snapshot.samples[last_sample_count..]; + last_sample_count = snapshot.samples.len(); + stream + .push_audio(delta, snapshot.sample_rate) + .map_err(|e| e.to_string())?; + for step in stream.drain_ready().map_err(|e| e.to_string())? { + if !step.text.trim().is_empty() { + last_text = step.text; + last_tokens = step.tokens.len(); + print_streaming_transcript(&last_text); + } + if step.reached_eos { + break; + } + } + Ok(()) + }); + + let (samples, sample_rate) = match record_result { + Ok(recording) => recording, + Err(e) => { + eprintln!("Recording failed: {e}"); + std::process::exit(1); + } + }; + + if samples.len() > last_sample_count { + let delta = &samples[last_sample_count..]; + if let Err(e) = stream.push_audio(delta, sample_rate) { + eprintln!("Streaming transcription failed: {e}"); + std::process::exit(1); + } + } + stream.finish(); + match stream.drain_ready() { + Ok(steps) => { + for step in steps { + if !step.text.trim().is_empty() { + last_text = step.text; + last_tokens = step.tokens.len(); + print_streaming_transcript(&last_text); + } + if step.reached_eos { + break; + } + } + } + Err(e) => { + eprintln!("Streaming transcription failed: {e}"); + std::process::exit(1); + } + } + + INTERRUPTED.store(false, Ordering::Relaxed); + + if !last_text.trim().is_empty() { + println!(); + if !QUIET.load(Ordering::Relaxed) { + let _ = io::stderr().flush(); + eprintln!("\n({last_tokens} tokens)"); + } + } else { + println!(); + eprintln!("No speech detected in recording."); + } +} + +fn print_streaming_transcript(text: &str) { + print!("\r\x1b[2K{}", text.trim()); + let _ = io::stdout().flush(); +} + /// Record from mic (VAD auto-stop), transcribe, return result. /// /// Used by the JSON-RPC `listen` method. Returns `None` if no speech From dc675158eb691d205dced2259dbafda5041cade1 Mon Sep 17 00:00:00 2001 From: Kyle Kelley Date: Tue, 7 Jul 2026 10:06:11 -0700 Subject: [PATCH 5/5] fix(voxtral): align native stream prefixes --- crates/voice-stt/src/lib.rs | 57 ++++++++++++++++++++ crates/voice-voxtral/src/realtime_session.rs | 56 +++++++++---------- crates/voice-voxtral/src/realtime_stream.rs | 2 + 3 files changed, 88 insertions(+), 27 deletions(-) diff --git a/crates/voice-stt/src/lib.rs b/crates/voice-stt/src/lib.rs index 0004104..0326a83 100644 --- a/crates/voice-stt/src/lib.rs +++ b/crates/voice-stt/src/lib.rs @@ -903,6 +903,63 @@ mod tests { assert_eq!(result.sample_rate, 16_000); } + #[test] + #[ignore = "loads Voxtral Realtime weights and runs native streaming inference"] + fn voxtral_native_stream_smoke_transcribes_eval_recording() { + let device = default_stt_device().unwrap(); + let mut model = load_backend_model_on_device( + SttBackend::Voxtral, + default_model_for_backend(SttBackend::Voxtral), + device, + ) + .unwrap(); + model.set_max_new_tokens(128); + + let audio_path = Path::new(env!("CARGO_MANIFEST_DIR")) + .join("../../eval/recordings/002.wav") + .canonicalize() + .unwrap(); + let audio = load_audio_file(&audio_path).unwrap(); + let mut stream = model.stream_session().unwrap(); + let chunk_samples = ((audio.sample_rate as usize) * 80 / 1_000).max(1); + let mut latest = None; + + for chunk in audio.samples.chunks(chunk_samples) { + stream.push_audio(chunk, audio.sample_rate).unwrap(); + for step in stream.drain_ready().unwrap() { + if !step.text.trim().is_empty() { + latest = Some(step.text); + } + if step.reached_eos { + break; + } + } + } + + stream.finish(); + for step in stream.drain_ready().unwrap() { + if !step.text.trim().is_empty() { + latest = Some(step.text); + } + if step.reached_eos { + break; + } + } + + let transcript = latest.unwrap_or_default(); + eprintln!("voxtral native stream transcript: {transcript}"); + assert!( + !transcript.trim().is_empty(), + "Voxtral native stream produced no transcript for {}", + audio_path.display() + ); + assert!( + transcript.to_ascii_lowercase().contains("seashore"), + "unexpected Voxtral native stream transcript for {}: {transcript:?}", + audio_path.display() + ); + } + #[test] fn test_resample_identity() { let sr = 16000u32; diff --git a/crates/voice-voxtral/src/realtime_session.rs b/crates/voice-voxtral/src/realtime_session.rs index 7d9dc22..e3b697e 100644 --- a/crates/voice-voxtral/src/realtime_session.rs +++ b/crates/voice-voxtral/src/realtime_session.rs @@ -22,15 +22,15 @@ pub struct VoxtralRealtimeStreamStep { /// Stateful Voxtral Realtime decoding session. /// /// This implements the model-native token/audio feedback loop. It intentionally -/// does not cache decoder KV state yet; each step recomputes the current text -/// history while preserving the externally visible streaming contract. +/// does not cache decoder KV state or incremental audio encoder state yet; each +/// step recomputes the current prefix while preserving the externally visible +/// streaming contract. pub struct VoxtralRealtimeStreamSession<'a> { transcriber: &'a VoxtralRealtimeTranscriber, buffer: VoxtralRealtimeStreamBuffer, options: VoxtralRealtimeTranscriptionOptions, input_token_ids: Vec, output_tokens: Vec, - audio_embeddings: Option, done: bool, } @@ -56,7 +56,6 @@ impl<'a> VoxtralRealtimeStreamSession<'a> { options, input_token_ids: Vec::new(), output_tokens: Vec::new(), - audio_embeddings: None, done: false, }) } @@ -92,26 +91,15 @@ impl<'a> VoxtralRealtimeStreamSession<'a> { return Ok(None); }; - let window_embeddings = self.transcriber.encode_stream_window_embeddings(&window)?; - self.audio_embeddings = Some(match &self.audio_embeddings { - Some(existing) => { - Tensor::cat(&[existing, &window_embeddings], 1).map_err(candle_err)? - } - None => window_embeddings, - }); self.input_token_ids .extend(window.input_token_ids.iter().copied()); - - let audio_embeddings = self - .audio_embeddings - .as_ref() - .expect("audio embeddings are set above"); + let audio_embeddings = self.transcriber.encode_stream_prefix_embeddings(&window)?; let generation = self .transcriber .text_decoder .greedy_decode_audio_embeddings_with_prompt( &self.transcriber.token_embeddings, - audio_embeddings, + &audio_embeddings, &self.input_token_ids, self.options.delay_tokens, 1, @@ -149,7 +137,28 @@ impl VoxtralRealtimeTranscriber { &self, window: &VoxtralRealtimeStreamWindow, ) -> Result { - let mel = realtime_log_mel_spectrogram(&self.config, &window.audio_samples)?; + self.encode_stream_embeddings_for_samples( + &window.audio_samples, + window.input_token_ids.len(), + ) + } + + pub fn encode_stream_prefix_embeddings( + &self, + window: &VoxtralRealtimeStreamWindow, + ) -> Result { + self.encode_stream_embeddings_for_samples( + &window.prefix_audio_samples, + window.input_token_end, + ) + } + + fn encode_stream_embeddings_for_samples( + &self, + samples: &[f32], + expected_tokens: usize, + ) -> Result { + let mel = realtime_log_mel_spectrogram(&self.config, samples)?; let input_features = Tensor::from_vec( mel.to_channel_major(), (1, mel.mel_bins, mel.frames), @@ -158,16 +167,11 @@ impl VoxtralRealtimeTranscriber { .map_err(candle_err)? .to_dtype(self.token_embeddings.tok_embeddings.embeddings().dtype()) .map_err(candle_err)?; - let audio_start_pos = window - .input_token_start - .checked_mul(self.config.downsample_factor()) - .ok_or_else(|| VoxtralError::InvalidConfig("stream audio position overflow".into()))?; let embeddings = self .audio_modules - .forward(&input_features, audio_start_pos) + .forward(&input_features, 0) .map_err(candle_err)?; let actual_tokens = embeddings.dim(1).map_err(candle_err)?; - let expected_tokens = window.input_token_ids.len(); if actual_tokens < expected_tokens { return Err(VoxtralError::Candle(format!( "stream window produced {actual_tokens} audio embeddings for {expected_tokens} input tokens" @@ -176,9 +180,7 @@ impl VoxtralRealtimeTranscriber { if actual_tokens == expected_tokens { return Ok(embeddings); } - embeddings - .narrow(1, actual_tokens - expected_tokens, expected_tokens) - .map_err(candle_err) + embeddings.narrow(1, 0, expected_tokens).map_err(candle_err) } } diff --git a/crates/voice-voxtral/src/realtime_stream.rs b/crates/voice-voxtral/src/realtime_stream.rs index dfc04c9..d00910f 100644 --- a/crates/voice-voxtral/src/realtime_stream.rs +++ b/crates/voice-voxtral/src/realtime_stream.rs @@ -28,6 +28,7 @@ pub struct VoxtralRealtimeStreamWindow { pub sequence: usize, pub input_token_ids: Vec, pub audio_samples: Vec, + pub prefix_audio_samples: Vec, pub frame_start_sample: usize, pub frame_end_sample: usize, pub stride_start_sample: usize, @@ -232,6 +233,7 @@ impl VoxtralRealtimeStreamBuffer { sequence: self.sequence, input_token_ids, audio_samples: self.samples[frame_start..frame_end].to_vec(), + prefix_audio_samples: self.samples[..frame_end].to_vec(), frame_start_sample: frame_start, frame_end_sample: frame_end, stride_start_sample: self.stride_start_sample,