From 1cf075fd1dac93cbe308e849e150200e9e856168 Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Fri, 17 Apr 2026 13:21:51 +0800
Subject: [PATCH] Qwen3ASR add vad data recognition
---
README.md | 17 +++++-------
README.zh-CN.md | 24 +++++++----------
docs/changelog.md | 6 +++++
docs/changelog.zh-CN.md | 6 +++++
src/models/common/generate.rs | 26 ++++++++++++++++++
src/models/common/modules.rs | 33 +++++++++++++++++++++++
src/models/fire_red_vad/vad.rs | 22 +++++++--------
src/models/qwen3_5/model.rs | 2 +-
src/models/qwen3_asr/generate.rs | 45 +++++++++++++++++++++++++++++--
src/models/qwen3_asr/model.rs | 42 ++++++++++++++++++++++++++---
src/models/qwen3_asr/processor.rs | 43 ++++++++++++++++++++++++++---
tests/test_fun_asr_nano.rs | 2 +-
tests/test_glm_asr_nano.rs | 2 +-
tests/test_qwen3_asr.rs | 4 +--
14 files changed, 221 insertions(+), 53 deletions(-)
diff --git a/README.md b/README.md
index c761e8e..0419534 100644
--- a/README.md
+++ b/README.md
@@ -49,6 +49,12 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
## Changelog
+### 2026-04-17
+- Qwen3ASR add vad data recognition
+
+### 2026-04-16
+- fix FireRedVAD fsmn cache bug
+
### 2026-04-15
- add FireRedVAD
@@ -71,17 +77,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
- \...\ The content of the thought chain is returned using the reasoning_content field.
- chat response add time info
-### 2026-04-01
-- refactor deepseek_ocr/fun_asr_nano generate code
-
-### 2026-03-31
-- add server and cli mod
-- aha model name use modelscope id replace
-- update WhichModel
-- Usage add time info
-- dependencies delete aha_openai_dive,chrono
-
-
**[View full changelog](docs/changelog.md)** →
diff --git a/README.zh-CN.md b/README.zh-CN.md
index 3fa75f7..6041b91 100644
--- a/README.zh-CN.md
+++ b/README.zh-CN.md
@@ -47,16 +47,21 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
## 更新日志
+### 2026-04-17
+- Qwen3ASR 增加 vad 数据识别
+
+### 2026-04-16
+- 修复 FireRedVAD fsmn 缓存问题
+
+### 2026-04-15
+- 添加 FireRedVAD
+
### 2026-04-10
- 修复 LiquidAI/LFM2.5-VL-450M chat_template 加载bug
### 2026-04-08
- 添加 VoxCPM2
-## Changelog
-### 2026-04-15
-- 添加 FireRedVAD
-
### 0.2.5 (2026-04-06)
- 添加 qwen3-embedding/qwen3-reranker/all-minilm-l6-v2
@@ -70,17 +75,6 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
- \...\ 思维链内容使用reasoning_content字段返回。
- 对话返回添加耗时信息
-### 2026-04-01
-- 重构 deepseek_ocr/fun_asr_nano 生成代码
-
-### 2026-03-31
-- 新增 server 和 cli 模块
-- aha模型名称使用 modelscope id 替换
-- 更新 WhichModel 枚举
-- Usage 增加时间信息
-- 删除 aha_openai_dive, chrono 依赖
-
-
**[查看完整更新日志](docs/changelog.zh-CN.md)** →
diff --git a/docs/changelog.md b/docs/changelog.md
index 144a3ca..a624ea8 100644
--- a/docs/changelog.md
+++ b/docs/changelog.md
@@ -5,6 +5,12 @@ All notable changes to aha will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
+### 2026-04-17
+- Qwen3ASR add vad data recognition
+
+### 2026-04-16
+- fix FireRedVAD fsmn cache bug
+
### 2026-04-15
- add FireRedVAD
diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md
index 8520fef..eeab560 100644
--- a/docs/changelog.zh-CN.md
+++ b/docs/changelog.zh-CN.md
@@ -5,6 +5,12 @@
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
+### 2026-04-17
+- Qwen3ASR 增加 vad 数据识别
+
+### 2026-04-16
+- 修复 FireRedVAD fsmn 缓存问题
+
### 2026-04-15
- 添加 FireRedVAD
diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs
index b80083f..6563aa2 100644
--- a/src/models/common/generate.rs
+++ b/src/models/common/generate.rs
@@ -117,6 +117,32 @@ fn sample_and_push(
generated.push(token);
Ok(token)
}
+pub fn generate_generic_text(
+ model: &mut M,
+ tokenizer: &TokenizerModel,
+ input_ids: Tensor,
+ data: MultiModalData,
+ ctx: &mut GenerationContext,
+) -> Result {
+ let mut generated = Vec::new();
+ let eos_ids = model.stop_token_ids();
+ let logits = model.forward_initial(&input_ids, ctx.seqlen_offset, data)?;
+ let next_token = sample_and_push(ctx, &logits, &mut generated)?;
+ let mut input_ids = ctx.prepare_for_next_token(next_token)?;
+
+ // 自回归循环
+ for _ in 1..ctx.sample_len {
+ let logits = model.forward_step(&input_ids, ctx.seqlen_offset)?;
+ let next_token = sample_and_push(ctx, &logits, &mut generated)?;
+
+ if eos_ids.contains(&next_token) {
+ break;
+ }
+ input_ids = ctx.prepare_for_next_token(next_token)?;
+ }
+ let text = tokenizer.token_decode(generated)?;
+ Ok(text)
+}
pub fn generate_generic(
model: &mut M,
diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs
index da76a35..cc0ef8a 100644
--- a/src/models/common/modules.rs
+++ b/src/models/common/modules.rs
@@ -12,6 +12,39 @@ use crate::{
utils::tensor_utils::{pad_replicate_last_dim, prepare_causal_attention_mask, repeat_kv},
};
+#[derive(Debug)]
+pub struct VadFrameResult {
+ pub is_speech: bool,
+ pub is_speech_start: bool,
+ pub is_i16: bool,
+ pub orig_audio: Option,
+ pub kaldi_audio: Option,
+ pub model_name: String,
+ pub mode: String,
+}
+
+#[derive(Debug)]
+pub struct AsrResult {
+ pub is_empty: bool,
+ pub text: Option,
+}
+
+impl AsrResult {
+ pub fn init_empty() -> Self {
+ AsrResult {
+ is_empty: true,
+ text: None,
+ }
+ }
+
+ pub fn init(text: String) -> Self {
+ AsrResult {
+ is_empty: false,
+ text: Some(text),
+ }
+ }
+}
+
#[derive(Debug, Clone)]
pub struct GateUpDownMLP {
gate_proj: Linear,
diff --git a/src/models/fire_red_vad/vad.rs b/src/models/fire_red_vad/vad.rs
index 45ba9eb..c90947e 100644
--- a/src/models/fire_red_vad/vad.rs
+++ b/src/models/fire_red_vad/vad.rs
@@ -3,10 +3,13 @@ use candle_core::{D, DType, Device, Tensor};
use candle_nn::VarBuilder;
use crate::{
- models::fire_red_vad::{
- config::{DetectModelConfig, FireRedVadConfig},
- model::DetectModel,
- processor::{AudioFeat, VadPostprocessor},
+ models::{
+ common::modules::VadFrameResult,
+ fire_red_vad::{
+ config::{DetectModelConfig, FireRedVadConfig},
+ model::DetectModel,
+ processor::{AudioFeat, VadPostprocessor},
+ },
},
utils::{
audio_utils::{resample_audio_from_bytes, resample_audio_from_vec_f32},
@@ -23,15 +26,6 @@ pub struct VadResult {
pub mode: String,
}
-#[derive(Debug)]
-pub struct VadFrameResult {
- pub is_speech: bool,
- pub orig_audio: Option,
- pub kaldi_audio: Option,
- pub model_name: String,
- pub mode: String,
-}
-
pub struct FireRedVad {
audio_feat: AudioFeat,
vad_model: DetectModel,
@@ -107,6 +101,8 @@ impl FireRedVad {
if preds_sum as f32 > probs.dim(0)? as f32 * self.cfg.speech_threshold {
Ok(Some(VadFrameResult {
is_speech: true,
+ is_i16: true,
+ is_speech_start: false, // TODO: is start speech, asr to clear cache
orig_audio: Some(audio_frame.clone()),
kaldi_audio: Some(feats),
model_name: self.model_name.clone(),
diff --git a/src/models/qwen3_5/model.rs b/src/models/qwen3_5/model.rs
index 49429e0..be08d66 100644
--- a/src/models/qwen3_5/model.rs
+++ b/src/models/qwen3_5/model.rs
@@ -1455,7 +1455,7 @@ impl InferenceModel for Qwen3_5Model {
) -> Result {
if data.data_vec.len() != 4 {
return Err(anyhow::anyhow!(
- "Lfm2VL process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw"
+ "Qwen3.5 process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw"
));
}
let pixel_values = &data.data_vec[0];
diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs
index 019ec0a..2ee1e1e 100644
--- a/src/models/qwen3_asr/generate.rs
+++ b/src/models/qwen3_asr/generate.rs
@@ -1,7 +1,11 @@
use std::time::Instant;
use crate::{
- models::common::generate::get_logit_processor,
+ models::common::{
+ MultiModalData,
+ generate::{GenerationContext, generate_generic_text, get_logit_processor},
+ modules::{AsrResult, VadFrameResult},
+ },
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time},
};
@@ -39,6 +43,7 @@ pub struct Qwen3AsrGenerateModel<'a> {
eos_token_id2: u32,
generation_config: Qwen3ASRGenerationConfig,
model_name: String,
+ default_template: String,
}
impl<'a> Qwen3AsrGenerateModel<'a> {
@@ -59,7 +64,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
let dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
- let qwen3_asr = Qwen3ASRModel::new(vb, &cfg)?;
+ let qwen3_asr = Qwen3ASRModel::new(vb, &cfg, generation_config.eos_token_id.clone())?;
let model_name = std::path::Path::new(path)
.file_name()
.and_then(|s| s.to_str())
@@ -76,8 +81,44 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config,
model_name,
+ default_template: "<|im_start|>system\n<|im_end|>\n<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|><|im_end|>\n<|im_start|>assistant\n".to_string(),
})
}
+
+ pub fn audio_recognize(&mut self, vad_res: VadFrameResult) -> Result {
+ if !vad_res.is_speech || vad_res.orig_audio.is_none() {
+ return Ok(AsrResult::init_empty());
+ }
+ if vad_res.is_speech_start {
+ self.qwen3_asr.clear_kv_cache();
+ }
+ let audio_data =
+ self.processor
+ .process_vad_res(&self.default_template, vad_res, &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(
+ None,
+ None,
+ None,
+ None,
+ None,
+ 32432,
+ input_ids.dim(1)?,
+ 512,
+ self.device.clone(),
+ );
+ let data_vec = vec![input_features];
+ let data = MultiModalData::new(data_vec);
+ let text = generate_generic_text(
+ &mut self.qwen3_asr,
+ &self.tokenizer,
+ input_ids,
+ data,
+ &mut ctx,
+ )?;
+ Ok(AsrResult::init(text))
+ }
}
impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
diff --git a/src/models/qwen3_asr/model.rs b/src/models/qwen3_asr/model.rs
index 294400d..640151f 100644
--- a/src/models/qwen3_asr/model.rs
+++ b/src/models/qwen3_asr/model.rs
@@ -7,7 +7,10 @@ use candle_nn::{
use crate::{
models::{
- common::modules::{NaiveAttention, get_conv2d, get_layer_norm},
+ common::{
+ InferenceModel,
+ modules::{NaiveAttention, get_conv2d, get_layer_norm},
+ },
qwen3::model::Qwen3DecoderLayer,
qwen3_asr::{
config::{
@@ -364,12 +367,16 @@ impl Qwen3ASRThinker {
pub struct Qwen3ASRModel {
thinker: Qwen3ASRThinker,
+ stop_token_ids: Vec,
}
impl Qwen3ASRModel {
- pub fn new(vb: VarBuilder, config: &Qwen3ASRConfig) -> Result {
+ pub fn new(vb: VarBuilder, config: &Qwen3ASRConfig, eos_ids: Vec) -> Result {
let thinker = Qwen3ASRThinker::new(vb.pp("thinker"), &config.thinker_config)?;
- Ok(Self { thinker })
+ Ok(Self {
+ thinker,
+ stop_token_ids: eos_ids,
+ })
}
pub fn forward(
@@ -387,3 +394,32 @@ impl Qwen3ASRModel {
self.thinker.clear_kv_cache();
}
}
+
+impl InferenceModel for Qwen3ASRModel {
+ fn forward_initial(
+ &mut self,
+ input_ids: &Tensor,
+ seqlen_offset: usize,
+ data: crate::models::common::MultiModalData,
+ ) -> Result {
+ if data.data_vec.len() != 1 {
+ return Err(anyhow::anyhow!(
+ "Qwen3 asr process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw"
+ ));
+ }
+ let input_features = &data.data_vec[0];
+ self.forward(input_ids, seqlen_offset, input_features.as_ref())
+ }
+
+ fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result {
+ self.forward(input_ids, seqlen_offset, None)
+ }
+
+ fn clear_cache(&mut self) {
+ self.clear_kv_cache();
+ }
+
+ fn stop_token_ids(&self) -> Vec {
+ self.stop_token_ids.clone()
+ }
+}
diff --git a/src/models/qwen3_asr/processor.rs b/src/models/qwen3_asr/processor.rs
index 08499b8..39ca5c4 100644
--- a/src/models/qwen3_asr/processor.rs
+++ b/src/models/qwen3_asr/processor.rs
@@ -1,7 +1,8 @@
use crate::{
- models::common::modules::float_range_normalize, params::chat::ChatCompletionParameters,
+ models::common::modules::{VadFrameResult, float_range_normalize},
+ params::chat::ChatCompletionParameters,
};
-use anyhow::Result;
+use anyhow::{Result, anyhow};
use candle_core::{Device, Tensor};
use crate::{
@@ -81,7 +82,7 @@ impl Qwen3AsrProcessor {
})
}
- pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result> {
+ pub fn extract_audio_vec(&self, mes: &ChatCompletionParameters) -> Result> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
audio_tensors.iter().map(float_range_normalize).collect()
}
@@ -96,6 +97,40 @@ impl Qwen3AsrProcessor {
text.replace("<|audio_placeholder|>", &self.audio_token)
}
+ pub fn process_vad_res(
+ &self,
+ render: &str,
+ vad_res: VadFrameResult,
+ 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)
+ } else {
+ Err(anyhow!("vad_res orig_audio is none!"))
+ }
+ }
+
pub fn process_info(
&self,
mes: &ChatCompletionParameters,
@@ -122,7 +157,7 @@ impl Qwen3AsrProcessor {
render = format!("{}language {}''", render, lang);
}
}
- let audio_tensors = self.process_audio(mes)?;
+ let audio_tensors = self.extract_audio_vec(mes)?;
let audio_len = audio_tensors.len();
if audio_len != audio_count {
return Err(anyhow::anyhow!("audio_pad num != audio num"));
diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs
index c17cead..cfec80b 100644
--- a/tests/test_fun_asr_nano.rs
+++ b/tests/test_fun_asr_nano.rs
@@ -21,7 +21,7 @@ fn fun_asr_nano_generate() -> Result<()> {
"type": "audio",
"audio_url":
{
- "url": "file://./assets/audio/zh.mp3"
+ "url": "file://./assets/audio/voice_01.wav"
}
},
{
diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs
index 0d75e64..110e48a 100644
--- a/tests/test_glm_asr_nano.rs
+++ b/tests/test_glm_asr_nano.rs
@@ -22,7 +22,7 @@ fn glm_asr_nano_generate() -> Result<()> {
"type": "audio",
"audio_url":
{
- "url": "file://./assets/audio/zh.mp3"
+ "url": "file://./assets/audio/voice_01.wav"
}
},
{
diff --git a/tests/test_qwen3_asr.rs b/tests/test_qwen3_asr.rs
index 5e750c6..d2ce259 100644
--- a/tests/test_qwen3_asr.rs
+++ b/tests/test_qwen3_asr.rs
@@ -6,7 +6,7 @@ use anyhow::Result;
use rocket::futures::StreamExt;
#[test]
fn qwen3_asr_generate() -> Result<()> {
- // RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_generate -r -- --nocapture
+ // RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3_asr qwen3_asr_generate -r -- --nocapture
let save_dir =
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
let model_path = format!("{}/Qwen/Qwen3-ASR-0.6B/", save_dir); //Qwen/Qwen3-ASR-1.7B
@@ -21,7 +21,7 @@ fn qwen3_asr_generate() -> Result<()> {
"type": "audio",
"audio_url":
{
- "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav"
+ "url": "file://./assets/audio/voice_01.wav"
}
}
]