diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index 18810f5..7a66d9a 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -85,13 +85,24 @@ impl<'a> Qwen3AsrGenerateModel<'a> { }) } - pub fn audio_recognize(&mut self, vad_res: VadFrameResult) -> Result { + pub fn asr_vad_res(&mut self, vad_res: VadFrameResult) -> Result { if !vad_res.is_speech || vad_res.orig_audio.is_none() { return Ok(AsrResult::init_empty()); } - let audio_data = - self.processor - .process_vad_res(&self.default_template, vad_res, &self.tokenizer)?; + if let Some(audio) = vad_res.orig_audio { + self.asr_audio(&audio, vad_res.is_i16) + } else { + Ok(AsrResult::init_empty()) + } + } + + pub fn asr_audio(&mut self, audio: &Tensor, is_i16: bool) -> Result { + let audio_data = self.processor.process_audio_tensor( + &self.default_template, + audio, + is_i16, + &self.tokenizer, + )?; let input_ids = audio_data.input_ids.clone(); let input_features = Some(audio_data.input_features.clone().to_dtype(self.dtype)?); let mut ctx = GenerationContext::new( diff --git a/src/models/qwen3_asr/processor.rs b/src/models/qwen3_asr/processor.rs index 39ca5c4..24b225b 100644 --- a/src/models/qwen3_asr/processor.rs +++ b/src/models/qwen3_asr/processor.rs @@ -97,6 +97,37 @@ impl Qwen3AsrProcessor { text.replace("<|audio_placeholder|>", &self.audio_token) } + pub fn process_audio_tensor( + &self, + render: &str, + audio: &Tensor, + is_i16: bool, + tokenizer: &TokenizerModel, + ) -> Result { + let audio_len = audio.dim(0)? as f32; + if audio_len > self.sample_rate as f32 * self.max_asr_input_seconds { + return Err(anyhow!("vad_res orig_audio is too long!")); + } + let mut audio = audio.unsqueeze(0)?; + if is_i16 { + audio = audio.affine(1.0 / 32768.0, 0.0)?; + } + audio = float_range_normalize(&audio)?; + let (input_features, _) = + self.whisper_feature_extracor + .call(&audio, self.sample_rate, false)?; + let audio_len = input_features.dim(2)?; + let output_len = get_feat_extract_output_lengths(audio_len); + let text = self.replace_special_tokens(render, output_len); + let input_ids = tokenizer.text_encode(text, &self.device)?; + let input_features = input_features.squeeze(0)?; + let audio_data = AudioData { + input_features, + input_ids, + }; + Ok(audio_data) + } + pub fn process_vad_res( &self, render: &str, @@ -104,28 +135,7 @@ impl Qwen3AsrProcessor { tokenizer: &TokenizerModel, ) -> Result { if let Some(audio) = &vad_res.orig_audio { - let audio_len = audio.dim(0)? as f32; - if audio_len > self.sample_rate as f32 * self.max_asr_input_seconds { - return Err(anyhow!("vad_res orig_audio is too long!")); - } - let mut audio = audio.unsqueeze(0)?; - if vad_res.is_i16 { - audio = audio.affine(1.0 / 32768.0, 0.0)?; - } - audio = float_range_normalize(&audio)?; - let (input_features, _) = - self.whisper_feature_extracor - .call(&audio, self.sample_rate, false)?; - let audio_len = input_features.dim(2)?; - let output_len = get_feat_extract_output_lengths(audio_len); - let text = self.replace_special_tokens(render, output_len); - let input_ids = tokenizer.text_encode(text, &self.device)?; - let input_features = input_features.squeeze(0)?; - let audio_data = AudioData { - input_features, - input_ids, - }; - Ok(audio_data) + self.process_audio_tensor(render, audio, vad_res.is_i16, tokenizer) } else { Err(anyhow!("vad_res orig_audio is none!")) }