update asr audio tensor
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
@@ -97,6 +97,37 @@ impl Qwen3AsrProcessor {
|
|||||||
text.replace("<|audio_placeholder|>", &self.audio_token)
|
text.replace("<|audio_placeholder|>", &self.audio_token)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn process_audio_tensor(
|
||||||
|
&self,
|
||||||
|
render: &str,
|
||||||
|
audio: &Tensor,
|
||||||
|
is_i16: bool,
|
||||||
|
tokenizer: &TokenizerModel,
|
||||||
|
) -> Result<AudioData> {
|
||||||
|
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(
|
pub fn process_vad_res(
|
||||||
&self,
|
&self,
|
||||||
render: &str,
|
render: &str,
|
||||||
@@ -104,28 +135,7 @@ impl Qwen3AsrProcessor {
|
|||||||
tokenizer: &TokenizerModel,
|
tokenizer: &TokenizerModel,
|
||||||
) -> Result<AudioData> {
|
) -> Result<AudioData> {
|
||||||
if let Some(audio) = &vad_res.orig_audio {
|
if let Some(audio) = &vad_res.orig_audio {
|
||||||
let audio_len = audio.dim(0)? as f32;
|
self.process_audio_tensor(render, audio, vad_res.is_i16, tokenizer)
|
||||||
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)
|
|
||||||
} else {
|
} else {
|
||||||
Err(anyhow!("vad_res orig_audio is none!"))
|
Err(anyhow!("vad_res orig_audio is none!"))
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user