update asr audio tensor

This commit is contained in:
jhqxxx
2026-04-17 20:24:35 +08:00
parent 0cb979145f
commit 3596b3ff84
2 changed files with 47 additions and 26 deletions
+15 -4
View File
@@ -85,13 +85,24 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
}) })
} }
pub fn audio_recognize(&mut self, vad_res: VadFrameResult) -> Result<AsrResult> { pub fn asr_vad_res(&mut self, vad_res: VadFrameResult) -> Result<AsrResult> {
if !vad_res.is_speech || vad_res.orig_audio.is_none() { if !vad_res.is_speech || vad_res.orig_audio.is_none() {
return Ok(AsrResult::init_empty()); return Ok(AsrResult::init_empty());
} }
let audio_data = if let Some(audio) = vad_res.orig_audio {
self.processor self.asr_audio(&audio, vad_res.is_i16)
.process_vad_res(&self.default_template, vad_res, &self.tokenizer)?; } else {
Ok(AsrResult::init_empty())
}
}
pub fn asr_audio(&mut self, audio: &Tensor, is_i16: bool) -> Result<AsrResult> {
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_ids = audio_data.input_ids.clone();
let input_features = Some(audio_data.input_features.clone().to_dtype(self.dtype)?); let input_features = Some(audio_data.input_features.clone().to_dtype(self.dtype)?);
let mut ctx = GenerationContext::new( let mut ctx = GenerationContext::new(
+14 -4
View File
@@ -97,19 +97,19 @@ impl Qwen3AsrProcessor {
text.replace("<|audio_placeholder|>", &self.audio_token) text.replace("<|audio_placeholder|>", &self.audio_token)
} }
pub fn process_vad_res( pub fn process_audio_tensor(
&self, &self,
render: &str, render: &str,
vad_res: VadFrameResult, audio: &Tensor,
is_i16: bool,
tokenizer: &TokenizerModel, tokenizer: &TokenizerModel,
) -> Result<AudioData> { ) -> Result<AudioData> {
if let Some(audio) = &vad_res.orig_audio {
let audio_len = audio.dim(0)? as f32; let audio_len = audio.dim(0)? as f32;
if audio_len > self.sample_rate as f32 * self.max_asr_input_seconds { if audio_len > self.sample_rate as f32 * self.max_asr_input_seconds {
return Err(anyhow!("vad_res orig_audio is too long!")); return Err(anyhow!("vad_res orig_audio is too long!"));
} }
let mut audio = audio.unsqueeze(0)?; let mut audio = audio.unsqueeze(0)?;
if vad_res.is_i16 { if is_i16 {
audio = audio.affine(1.0 / 32768.0, 0.0)?; audio = audio.affine(1.0 / 32768.0, 0.0)?;
} }
audio = float_range_normalize(&audio)?; audio = float_range_normalize(&audio)?;
@@ -126,6 +126,16 @@ impl Qwen3AsrProcessor {
input_ids, input_ids,
}; };
Ok(audio_data) Ok(audio_data)
}
pub fn process_vad_res(
&self,
render: &str,
vad_res: VadFrameResult,
tokenizer: &TokenizerModel,
) -> Result<AudioData> {
if let Some(audio) = &vad_res.orig_audio {
self.process_audio_tensor(render, audio, vad_res.is_i16, tokenizer)
} else { } else {
Err(anyhow!("vad_res orig_audio is none!")) Err(anyhow!("vad_res orig_audio is none!"))
} }