From f901f4cbfeeda927abc09f626e4856e0e0894d6d Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Wed, 15 Apr 2026 17:42:41 +0800 Subject: [PATCH] add FireRedVAD --- Cargo.lock | 32 ++- Cargo.toml | 1 + README.md | 3 + README.zh-CN.md | 3 + docs/changelog.md | 3 + docs/changelog.zh-CN.md | 3 + src/models/common/modules.rs | 38 ++++ src/models/fire_red_vad/config.rs | 136 +++++++++++++ src/models/fire_red_vad/mod.rs | 4 + src/models/fire_red_vad/model.rs | 291 +++++++++++++++++++++++++++ src/models/fire_red_vad/processor.rs | 244 ++++++++++++++++++++++ src/models/fire_red_vad/vad.rs | 161 +++++++++++++++ src/models/mod.rs | 2 + src/models/sam3/mod.rs | 1 + src/models/sam3/processor.rs | 0 src/models/voxcpm/model.rs | 14 +- src/utils/audio_utils.rs | 74 +++++-- src/utils/tensor_utils.rs | 31 +++ tests/messy_test.rs | 86 +++++++- tests/test_fire_red_vad.rs | 47 +++++ tests/test_robo_brain.rs | 48 ++++- tests/weight_test.rs | 59 ++++++ 22 files changed, 1244 insertions(+), 37 deletions(-) create mode 100644 src/models/fire_red_vad/config.rs create mode 100644 src/models/fire_red_vad/mod.rs create mode 100644 src/models/fire_red_vad/model.rs create mode 100644 src/models/fire_red_vad/processor.rs create mode 100644 src/models/fire_red_vad/vad.rs create mode 100644 src/models/sam3/mod.rs create mode 100644 src/models/sam3/processor.rs create mode 100644 tests/test_fire_red_vad.rs diff --git a/Cargo.lock b/Cargo.lock index 83aaf7c..64a2b97 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -38,6 +38,7 @@ dependencies = [ "half", "hound", "image", + "kaldi-native-fbank", "minijinja", "modelscope", "num", @@ -511,7 +512,7 @@ dependencies = [ "objc2-foundation", "objc2-metal", "rand 0.9.2", - "rand_distr", + "rand_distr 0.5.1", "rayon", "safetensors 0.7.0", "thiserror 2.0.18", @@ -1372,7 +1373,7 @@ dependencies = [ "half", "num-traits", "rand 0.9.2", - "rand_distr", + "rand_distr 0.5.1", ] [[package]] @@ -1886,7 +1887,7 @@ dependencies = [ "crunchy", "num-traits", "rand 0.9.2", - "rand_distr", + "rand_distr 0.5.1", "zerocopy", ] @@ -2397,6 +2398,21 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "kaldi-native-fbank" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f297ca19f4c2069ccc7969cba1d4c496b46c0b88f6ceddb88a9ae3200a6cf506" +dependencies = [ + "anyhow", + "log", + "rand 0.8.5", + "rand_distr 0.4.3", + "realfft", + "rustfft", + "thiserror 1.0.69", +] + [[package]] name = "lazy_static" version = "1.5.0" @@ -3422,6 +3438,16 @@ dependencies = [ "getrandom 0.3.4", ] +[[package]] +name = "rand_distr" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32cb0b9bc82b0a0876c2dd994a7e7a2683d3e7390ca40e6886785ef0c7e3ee31" +dependencies = [ + "num-traits", + "rand 0.8.5", +] + [[package]] name = "rand_distr" version = "0.5.1" diff --git a/Cargo.toml b/Cargo.toml index d9d1e23..e122bce 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -44,6 +44,7 @@ byteorder = "1.5.0" sentencepiece = "0.13.1" ahash = "0.8.12" derive_builder = "0.20.2" +kaldi-native-fbank = "0.1.0" [patch.crates-io] esaxx-rs = { path = "vendor/esaxx-rs" } diff --git a/README.md b/README.md index 0822e41..c761e8e 100644 --- a/README.md +++ b/README.md @@ -49,6 +49,9 @@ 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-15 +- add FireRedVAD + ### 2026-04-10 - fix LiquidAI/LFM2.5-VL-450M chat_template load bug diff --git a/README.zh-CN.md b/README.zh-CN.md index 4a3d0f6..3fa75f7 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -54,6 +54,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 - 添加 VoxCPM2 ## Changelog +### 2026-04-15 +- 添加 FireRedVAD + ### 0.2.5 (2026-04-06) - 添加 qwen3-embedding/qwen3-reranker/all-minilm-l6-v2 diff --git a/docs/changelog.md b/docs/changelog.md index 47f8aaf..144a3ca 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -5,6 +5,9 @@ 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-15 +- add FireRedVAD + ### 2026-04-10 - fix LiquidAI/LFM2.5-VL-450M chat_template load bug diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index 2c60aba..8520fef 100644 --- a/docs/changelog.zh-CN.md +++ b/docs/changelog.zh-CN.md @@ -5,6 +5,9 @@ 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/), 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 +### 2026-04-15 +- 添加 FireRedVAD + ### 2026-04-10 - 修复 LiquidAI/LFM2.5-VL-450M chat_template 加载bug diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs index a384551..da76a35 100644 --- a/src/models/common/modules.rs +++ b/src/models/common/modules.rs @@ -1335,6 +1335,44 @@ pub fn conv1d_depthwise(input: &Tensor, weight: &Tensor, bias: Option<&Tensor>) } } +pub fn native_conv1d(input: &Tensor, weight: &Tensor, mode: &str) -> Result { + // input: shape: (n) + // weight: shape: (kernel_size) + // mode: 'full', 'same', 'valid' + // full : out_len = n + kernel_size - 1 + // same : out_len = n + // valid: out_len = n - kernel_size + 1 + let kernel_size: usize = weight.dim(0)?; + let pad = match mode { + "full" => { + let pad = kernel_size - 1; + input.pad_with_zeros(0, pad, pad)? + } + "same" => { + let pad = (kernel_size - 1) / 2; + input.pad_with_zeros(0, pad, pad)? + } + "valid" => input.clone(), + _ => { + return Err(anyhow!( + "native_conv1d only support mode is 'full', 'same', 'valid'" + )); + } + }; + let in_len = pad.dim(0)?; + let len_out = in_len - kernel_size + 1; + let mut out = pad + .narrow(D::Minus1, 0, len_out)? + .broadcast_mul(&weight.narrow(D::Minus1, 0, 1)?)?; + for k in 1..kernel_size { + out = (out + + pad + .narrow(D::Minus1, k, len_out)? + .broadcast_mul(&weight.narrow(D::Minus1, k, 1)?)?)?; + } + Ok(out) +} + pub fn log10(t: &Tensor) -> Result { Ok(t.log()?.affine(1.0 / 10.0_f64.ln(), 0.0)?) } diff --git a/src/models/fire_red_vad/config.rs b/src/models/fire_red_vad/config.rs new file mode 100644 index 0000000..f9b17e0 --- /dev/null +++ b/src/models/fire_red_vad/config.rs @@ -0,0 +1,136 @@ +pub struct FireRedVadConfig { + pub smooth_window_size: usize, + pub speech_threshold: f32, + pub singing_threshold: f32, + pub music_threshold: f32, + pub pad_start_frame: usize, + pub min_speech_frame: usize, + pub max_speech_frame: usize, + pub min_event_frame: usize, + pub max_event_frame: usize, + pub min_silence_frame: usize, + pub merge_silence_frame: usize, + pub extend_speech_frame: usize, + pub chunk_max_frame: usize, +} + +impl FireRedVadConfig { + pub fn default_vad() -> Self { + Self { + smooth_window_size: 5, + speech_threshold: 0.4, + singing_threshold: 0.5, + music_threshold: 0.5, + pad_start_frame: 5, + min_speech_frame: 20, + max_speech_frame: 2000, + min_event_frame: 20, + max_event_frame: 2000, + min_silence_frame: 20, + merge_silence_frame: 0, + extend_speech_frame: 0, + chunk_max_frame: 30000, + } + } + + pub fn default_stream_vad() -> Self { + Self { + smooth_window_size: 1, + speech_threshold: 0.5, + singing_threshold: 0.5, + music_threshold: 0.5, + pad_start_frame: 5, + min_speech_frame: 8, + max_speech_frame: 2000, + min_event_frame: 20, + max_event_frame: 2000, + min_silence_frame: 20, + merge_silence_frame: 0, + extend_speech_frame: 0, + chunk_max_frame: 30000, + } + } + + pub fn default_aed() -> Self { + Self { + smooth_window_size: 5, + speech_threshold: 0.4, + singing_threshold: 0.5, + music_threshold: 0.5, + pad_start_frame: 5, + min_speech_frame: 8, + max_speech_frame: 2000, + min_event_frame: 20, + max_event_frame: 2000, + min_silence_frame: 20, + merge_silence_frame: 0, + extend_speech_frame: 0, + chunk_max_frame: 30000, + } + } +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct CMVNData { + pub cmvn: Vec>, +} + +pub struct DetectModelConfig { + pub idim: usize, + pub r: usize, + pub m: usize, + pub h: usize, + pub p: usize, + pub n1: usize, + pub s1: usize, + pub n2: usize, + pub s2: usize, + pub odim: usize, +} + +impl DetectModelConfig { + pub fn default_vad() -> Self { + Self { + idim: 80, + r: 8, + m: 1, + h: 256, + p: 128, + n1: 20, + s1: 1, + n2: 20, + s2: 1, + odim: 1, + } + } + + pub fn default_stream_vad() -> Self { + Self { + idim: 80, + r: 8, + m: 1, + h: 256, + p: 128, + n1: 20, + s1: 1, + n2: 0, + s2: 1, + odim: 1, + } + } + + pub fn default_aed() -> Self { + Self { + idim: 80, + r: 8, + m: 1, + h: 256, + p: 128, + n1: 20, + s1: 1, + n2: 20, + s2: 1, + odim: 3, + } + } +} diff --git a/src/models/fire_red_vad/mod.rs b/src/models/fire_red_vad/mod.rs new file mode 100644 index 0000000..a746448 --- /dev/null +++ b/src/models/fire_red_vad/mod.rs @@ -0,0 +1,4 @@ +pub mod config; +pub mod model; +pub mod processor; +pub mod vad; diff --git a/src/models/fire_red_vad/model.rs b/src/models/fire_red_vad/model.rs new file mode 100644 index 0000000..3191b75 --- /dev/null +++ b/src/models/fire_red_vad/model.rs @@ -0,0 +1,291 @@ +use anyhow::Result; +use candle_core::{D, Tensor}; +use candle_nn::{Conv1d, Linear, Module, VarBuilder, linear, linear_no_bias, ops::sigmoid}; + +use crate::{ + models::{ + common::modules::{conv1d_depthwise, get_conv1d}, + fire_red_vad::config::DetectModelConfig, + }, + utils::tensor_utils::{get_mask_from_lengths, masked_fill_zeros}, +}; + +pub struct FSMN { + lookback_padding: usize, + lookback_filter: Conv1d, + lookahead_filter: Option, + n1: usize, + s1: usize, + n2: usize, + s2: usize, +} + +impl FSMN { + pub fn new( + vb: VarBuilder, + p: usize, + n1: usize, + s1: usize, + n2: usize, + s2: usize, + ) -> Result { + let lookback_padding = (n1 - 1) * s1; + let lookback_filter = get_conv1d( + vb.pp("lookback_filter"), + p, + p, + n1, + lookback_padding, + 1, + s1, + p, + false, + )?; + let lookahead_filter = if n2 > 0 { + Some(get_conv1d( + vb.pp("lookahead_filter"), + p, + p, + n2, + (n2 - 1) * s2, + 1, + s2, + p, + false, + )?) + } else { + None + }; + Ok(Self { + lookback_padding, + lookback_filter, + lookahead_filter, + n1, + s1, + n2, + s2, + }) + } + + pub fn forward( + &self, + inputs: &Tensor, + mask: Option<&Tensor>, + cache: Option<&Tensor>, + ) -> Result<(Tensor, Tensor)> { + let t = inputs.dim(1)?; + let inputs = if let Some(mask) = mask { + masked_fill_zeros(inputs, mask)? + } else { + inputs.clone() + }; + // [N, T, P] -> [N, P, T] + let residual = inputs.permute((0, 2, 1))?.contiguous()?; + let inputs = if let Some(cache) = cache { + Tensor::cat(&[cache, &residual], 2)? + } else { + residual.clone() + }; + let start = inputs.dim(D::Minus1)? - self.lookback_padding; + let new_cache = inputs.narrow(D::Minus1, start, self.lookback_padding)?; + // conv1d_depthwise仅支持dilation=1的情况 + let lookback = if self.s1 == 1 { + let inputs = + inputs.pad_with_zeros(D::Minus1, self.lookback_padding, self.lookback_padding)?; + conv1d_depthwise( + &inputs, + self.lookback_filter.weight(), + self.lookback_filter.bias(), + )? + } else { + self.lookback_filter.forward(&inputs)? + }; + let mut memory = if self.n1 > 1 { + let len = lookback.dim(D::Minus1)? - (self.n1 - 1) * self.s1; + let mut lookback = lookback.narrow(D::Minus1, 0, len)?; + if let Some(cache) = cache { + let start = cache.dim(2)?; + let len = lookback.dim(D::Minus1)? - start; + lookback = lookback.narrow(D::Minus1, start, len)?; + } + residual.add(&lookback)? + } else { + residual.add(&lookback)? + }; + if self.n2 > 0 + && t > 1 + && let Some(ahead_filter) = &self.lookahead_filter + { + let lookahead = if self.s2 == 1 { + let inputs = inputs.pad_with_zeros( + D::Minus1, + self.lookback_padding, + self.lookback_padding, + )?; + conv1d_depthwise(&inputs, ahead_filter.weight(), ahead_filter.bias())? + } else { + ahead_filter.forward(&inputs)? + }; + let start = self.n2 * self.s2; + let len = lookahead.dim(D::Minus1)? - start; + let lookahead = lookahead.narrow(D::Minus1, start, len)?; + let lookahead = lookahead.pad_with_zeros(D::Minus1, 0, self.s2)?; + memory = memory.add(&lookahead)?; + } + memory = memory.permute((0, 2, 1))?.contiguous()?; + if let Some(mask) = mask { + memory = masked_fill_zeros(&memory, mask)?; + } + Ok((memory, new_cache)) + } +} + +struct DFSMNBlock { + fc1: Linear, // linear + relu + fc2: Linear, + fsmn: FSMN, +} + +impl DFSMNBlock { + pub fn new( + vb: VarBuilder, + h: usize, + p: usize, + n1: usize, + s1: usize, + n2: usize, + s2: usize, + ) -> Result { + let fc1 = linear(p, h, vb.pp("fc1.0"))?; + let fc2 = linear_no_bias(h, p, vb.pp("fc2"))?; + let fsmn = FSMN::new(vb.pp("fsmn"), p, n1, s1, n2, s2)?; + Ok(Self { fc1, fc2, fsmn }) + } + + pub fn forward( + &self, + inputs: &Tensor, + mask: Option<&Tensor>, + cache: Option<&Tensor>, + ) -> Result<(Tensor, Tensor)> { + let residual = inputs.clone(); + let h = self.fc1.forward(inputs)?.relu()?; + let p = self.fc2.forward(&h)?; + let (memory, new_cache) = self.fsmn.forward(&p, mask, cache)?; + let output = memory.add(&residual)?; + Ok((output, new_cache)) + } +} + +#[allow(clippy::upper_case_acronyms)] +struct DFSMN { + fc1: Linear, // linear + relu + fc2: Linear, // linear + relu + fsmn1: FSMN, + fsmns: Vec, + dnns: Vec, // linear + relu +} + +impl DFSMN { + pub fn new( + vb: VarBuilder, + d: usize, + r: usize, + m: usize, + h: usize, + p: usize, + n1: usize, + s1: usize, + n2: usize, + s2: usize, + ) -> Result { + let fc1 = linear(d, h, vb.pp("fc1.0"))?; + let fc2 = linear(h, p, vb.pp("fc2.0"))?; + let fsmn1 = FSMN::new(vb.pp("fsmn1"), p, n1, s1, n2, s2)?; + let mut fsmns = vec![]; + let vb_fsmns = vb.pp("fsmns"); + for i in 0..(r - 1) { + let block = DFSMNBlock::new(vb_fsmns.pp(i), h, p, n1, s1, n2, s2)?; + fsmns.push(block); + } + let vb_dnns = vb.pp("dnns"); + let mut dnns = vec![]; + for i in 0..m { + let in_dim = if i == 0 { p } else { h }; + let dnn = linear(in_dim, h, vb_dnns.pp(i))?; + dnns.push(dnn); + } + Ok(Self { + fc1, + fc2, + fsmn1, + fsmns, + dnns, + }) + } + + pub fn forward( + &self, + inputs: &Tensor, + input_lengths: Option<&Tensor>, + caches: Option<&Vec>, + ) -> Result<(Tensor, Vec)> { + let mask = if let Some(input_lengths) = input_lengths { + Some(get_mask_from_lengths(input_lengths)?) + } else { + None + }; + let h = self.fc1.forward(inputs)?.relu()?; + let p = self.fc2.forward(&h)?.relu()?; + let mut new_caches = vec![]; + let cache = caches.map(|caches| &caches[0]); + let (mut memory, mut new_cache) = self.fsmn1.forward(&p, mask.as_ref(), cache)?; + new_caches.push(new_cache); + let mut i = 1; + for fsmn in &self.fsmns { + let cache = caches.map(|caches| &caches[i]); + (memory, new_cache) = fsmn.forward(&memory, mask.as_ref(), cache)?; + new_caches.push(new_cache); + i += 1; + } + for dnn in &self.dnns { + memory = dnn.forward(&memory)?.relu()?; + } + Ok((memory, new_caches)) + } +} + +pub struct DetectModel { + dfsmn: DFSMN, + out: Linear, +} + +impl DetectModel { + pub fn new(vb: VarBuilder, cfg: DetectModelConfig) -> Result { + let dfsmn = DFSMN::new( + vb.pp("dfsmn"), + cfg.idim, + cfg.r, + cfg.m, + cfg.h, + cfg.p, + cfg.n1, + cfg.s1, + cfg.n2, + cfg.s2, + )?; + let out = linear(cfg.h, cfg.odim, vb.pp("out"))?; + Ok(Self { dfsmn, out }) + } + + pub fn forward( + &self, + feat: &Tensor, + caches: Option<&Vec>, + ) -> Result<(Tensor, Vec)> { + let (x, new_caches) = self.dfsmn.forward(feat, None, caches)?; + let logits = self.out.forward(&x)?; + let probs = sigmoid(&logits)?; + Ok((probs, new_caches)) + } +} diff --git a/src/models/fire_red_vad/processor.rs b/src/models/fire_red_vad/processor.rs new file mode 100644 index 0000000..48b436e --- /dev/null +++ b/src/models/fire_red_vad/processor.rs @@ -0,0 +1,244 @@ +use std::cmp::min; + +use anyhow::{Result, anyhow}; +use candle_core::{D, Device, IndexOp, Tensor}; +use kaldi_native_fbank::{ + FbankComputer, FbankOptions, + window::{Window, extract_window}, +}; + +use crate::{ + models::{ + common::modules::native_conv1d, + fire_red_vad::config::{CMVNData, FireRedVadConfig}, + }, + utils::{audio_utils::load_audio_with_resample, tensor_utils::apply_threshold}, +}; +pub struct CMVN { + dim: usize, + means: Tensor, + inverse_std_variances: Tensor, +} + +impl CMVN { + pub fn new(path: &str, device: &Device) -> Result { + let cmvn_path = path.to_string() + "/cmvn.json"; + assert!( + std::path::Path::new(&cmvn_path).exists(), + "cmvn path file not exists" + ); + let cmvn: CMVNData = serde_json::from_slice(&std::fs::read(cmvn_path)?)?; + let cmvn_data = Tensor::new(cmvn.cmvn, device)?; + assert_eq!(cmvn_data.rank(), 2); + assert_eq!(cmvn_data.dim(0)?, 2); + let (_, dim) = cmvn_data.dims2()?; + let dim = dim - 1; + let count = cmvn_data.i((0, dim))?.to_scalar::()?; + assert!(count >= 1.0); + let floor = 1e-20f32; + let means = cmvn_data.i((0, 0..dim))?.affine(1.0 / count as f64, 0.0)?; + let variance = cmvn_data + .i((1, 0..dim))? + .affine(1.0 / count as f64, 0.0)? + .sub(&means.powf(2.0)?)? + .clamp(floor, f32::MAX)?; + let inverse_std_variances = (1.0 / variance.sqrt()?)?; + Ok(Self { + dim, + means, + inverse_std_variances, + }) + } + + pub fn call(&self, xs: &Tensor) -> Result { + assert_eq!(xs.dim(D::Minus1)?, self.dim, "CMVN dim mismatch"); + let xs = xs.broadcast_sub(&self.means)?; + let xs = xs.broadcast_mul(&self.inverse_std_variances)?; + Ok(xs) + } +} + +pub struct KaldifeatFbank { + opts: FbankOptions, + win: Window, +} + +impl KaldifeatFbank { + pub fn new(num_mel_bins: usize, dither: f32) -> Result { + let mut opts = FbankOptions::default(); + opts.frame_opts.samp_freq = 16000.0; + opts.frame_opts.frame_length_ms = 25.0; + opts.frame_opts.frame_shift_ms = 10.0; + opts.frame_opts.dither = dither; + opts.frame_opts.snip_edges = true; + opts.mel_opts.num_bins = num_mel_bins; + opts.mel_opts.debug_mel = false; + opts.use_energy = false; + let win = Window::new(&opts.frame_opts) + .ok_or("window new error") + .map_err(|e| anyhow!("fbank comput err: {e}"))?; + Ok(Self { opts, win }) + } + + pub fn call(&self, wav_tensor: &Tensor) -> Result { + let mut comp = + FbankComputer::new(self.opts.clone()).map_err(|e| anyhow!("fbank comput err: {e}"))?; + + let padded = self.opts.frame_opts.padded_window_size(); + let wave = if wav_tensor.rank() == 1 { + wav_tensor.to_vec1::()? + } else if wav_tensor.rank() == 2 { + wav_tensor.squeeze(0)?.to_vec1::()? + } else { + return Err(anyhow!("not support wav dim: {}", wav_tensor.rank())); + }; + let mut feats = vec![]; + let mut window_buf = vec![0.0; padded]; + for frame in 0..230 { + let raw_log_energy = extract_window( + 0, + &wave, + frame, + &self.opts.frame_opts, + Some(&self.win), + &mut window_buf, + ) + .unwrap(); + let mut feat = vec![0.0; comp.dim()]; + comp.compute(raw_log_energy, 1.0, &mut window_buf, &mut feat); + feats.push(feat); + } + let feats = Tensor::new(feats, wav_tensor.device())?; + Ok(feats) + } +} + +pub struct AudioFeat { + cmvn: CMVN, + fbank: KaldifeatFbank, +} + +impl AudioFeat { + pub fn new(path: &str, device: &Device) -> Result { + let cmvn = CMVN::new(path, device)?; + let fbank = KaldifeatFbank::new(80, 0.0)?; + Ok(Self { cmvn, fbank }) + } + + pub fn extract(&self, wave_tensor: &Tensor) -> Result { + let fbank = self.fbank.call(wave_tensor)?; + let fbank = self.cmvn.call(&fbank)?; + Ok(fbank) + } + + pub fn extract_file(&self, audio_path: &str, device: &Device) -> Result<(Tensor, f32)> { + let wave_tensor = + load_audio_with_resample(audio_path, device, Some(16000), true)?.squeeze(0)?; + let dur = wave_tensor.dim(0)? as f32 / 16000.0; + let fbank = self.extract(&wave_tensor)?; + Ok((fbank, dur)) + } +} + +pub enum VadState { + SILENCE, + POSSIBLESPEECH, + SPEECH, + POSSIBLESILENCE, +} + +pub struct VadPostprocessor { + pub smooth_window_size: usize, + pub prob_threshold: f32, + pub pad_start_frame: usize, + pub min_speech_frame: usize, + pub max_speech_frame: usize, + pub min_silence_frame: usize, + pub merge_silence_frame: usize, + pub extend_speech_frame: usize, + pub frame_shift_s: f32, + pub frame_cnt: usize, + pub state: VadState, +} + +impl VadPostprocessor { + pub fn new(cfg: &FireRedVadConfig) -> Self { + Self { + smooth_window_size: cfg.smooth_window_size, + prob_threshold: cfg.speech_threshold, + pad_start_frame: cfg.pad_start_frame, + min_speech_frame: cfg.min_speech_frame, + max_speech_frame: cfg.max_speech_frame, + min_silence_frame: cfg.min_silence_frame, + merge_silence_frame: cfg.merge_silence_frame, + extend_speech_frame: cfg.extend_speech_frame, + frame_shift_s: 0.01, + frame_cnt: 0, + state: VadState::SILENCE, + } + } + + pub fn reset(&mut self) { + self.frame_cnt = 0; + } + + pub fn process_one(&self, probs: f32) -> Result { + // TODO: 状态管理 + let is_speech = probs >= self.prob_threshold; + Ok(is_speech) + } + + pub fn process_thresh(&self, raw_probs: &Tensor) -> Result { + let smoothed_probs = self.smooth_prob(raw_probs)?; + let binary_preds = apply_threshold(&smoothed_probs, self.prob_threshold)?; + Ok(binary_preds) + } + + pub fn process(&self, raw_probs: &Tensor, dur: f32) -> Result> { + let binary_preds = self.process_thresh(raw_probs)?; + self.decision_to_segment(&binary_preds, dur) + } + + pub fn decision_to_segment(&self, decisions: &Tensor, dur: f32) -> Result> { + let mut segments = vec![]; + let mut speech_start = -1; + let decisions = decisions.to_vec1::()?; + for (t, &flag) in decisions.iter().enumerate() { + if flag == 1 && speech_start == -1 { + speech_start = t as i32; + } else if flag == 0 && speech_start != -1 { + segments.push(( + speech_start as f32 * self.frame_shift_s, + t as f32 * self.frame_shift_s, + )); + speech_start = -1; + } + } + if speech_start != -1 { + let t = decisions.len() - 1; + let end_time = dur.min(t as f32 * self.frame_shift_s); + segments.push((speech_start as f32 * self.frame_shift_s, end_time)); + } + Ok(segments) + } + + fn smooth_prob(&self, probs: &Tensor) -> Result { + if self.smooth_window_size <= 1 { + Ok(probs.clone()) + } else { + let kernel_value = 1.0 / self.smooth_window_size as f32; + let probs_len = probs.dim(0)?; + let weight = Tensor::new(vec![kernel_value; self.smooth_window_size], probs.device())?; + let mut moothed = native_conv1d(probs, &weight, "full")?.i(0..probs_len)?; + let mean_len = min(self.smooth_window_size - 1, probs_len); + let mut mean_vec = vec![]; + for i in 0..mean_len { + let mean = probs.i(0..i + 1)?.mean(0)?.to_scalar::()?; + mean_vec.push(mean); + } + let means = Tensor::new(mean_vec, probs.device())?; + moothed = moothed.slice_assign(&[(0..mean_len)], &means)?; + Ok(moothed) + } + } +} diff --git a/src/models/fire_red_vad/vad.rs b/src/models/fire_red_vad/vad.rs new file mode 100644 index 0000000..0192e11 --- /dev/null +++ b/src/models/fire_red_vad/vad.rs @@ -0,0 +1,161 @@ +use anyhow::{Result, anyhow}; +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}, + }, + utils::{ + audio_utils::load_audio_with_resample_from_bytes, find_type_files, get_device, + tensor_utils::split_tensor_with_size, + }, +}; + +#[derive(Debug)] +pub struct VadResult { + pub dur: f32, + pub timestamps: Vec<(f32, f32)>, + pub model_name: String, + 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, + vad_postprocessor: VadPostprocessor, + model_name: String, + device: Device, + cfg: FireRedVadConfig, + caches: Option>, + frame_length_sample: usize, +} + +impl FireRedVad { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let device = get_device(device); + let audio_feat = AudioFeat::new(path, &device)?; + let model_list = find_type_files(path, "safetensors")?; + let dtype = dtype.unwrap_or(DType::F32); + let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("VAD") + .to_string(); + let (model_cfg, cfg) = if model_name.to_lowercase().contains("stream") { + ( + DetectModelConfig::default_stream_vad(), + FireRedVadConfig::default_stream_vad(), + ) + } else if model_name.to_lowercase().contains("aed") { + ( + DetectModelConfig::default_aed(), + FireRedVadConfig::default_aed(), + ) // TODO: aed + } else { + ( + DetectModelConfig::default_vad(), + FireRedVadConfig::default_vad(), + ) + }; + let vad_model = DetectModel::new(vb, model_cfg)?; + let vad_postprocessor = VadPostprocessor::new(&cfg); + Ok(Self { + audio_feat, + vad_model, + vad_postprocessor, + model_name, + device, + cfg, + caches: None, + frame_length_sample: 400, + }) + } + + pub fn detect_frame(&mut self, audio_bytes: Vec) -> Result> { + if !self.model_name.to_lowercase().contains("stream") { + return Err(anyhow!("only stream model support detect_frame")); + } + let audio_frame = + load_audio_with_resample_from_bytes(audio_bytes, &self.device, Some(16000), true)? + .squeeze(0)?; + if audio_frame.dim(0)? < self.frame_length_sample { + return Err(anyhow!( + "Expected {} samples, got {}", + self.frame_length_sample, + audio_frame.dim(0)? + )); + } + let feats = self.audio_feat.extract(&audio_frame)?; + let (probs, caches) = self + .vad_model + .forward(&feats.unsqueeze(0)?, self.caches.as_ref())?; + self.caches = Some(caches); + let probs = probs.squeeze(D::Minus1)?.squeeze(0)?; + let binary_preds = self + .vad_postprocessor + .process_thresh(&probs)? + .to_dtype(DType::U32)?; + let preds_sum = binary_preds.sum_all()?.to_scalar::()?; + if preds_sum as f32 > probs.dim(0)? as f32 * self.cfg.speech_threshold { + Ok(Some(VadFrameResult { + is_speech: true, + orig_audio: Some(audio_frame), + kaldi_audio: Some(feats), + model_name: self.model_name.clone(), + mode: "speech".to_string(), + })) + } else { + Ok(None) + } + } + + pub fn detect_file(&self, audio_path: &str) -> Result { + let (feats, dur) = self.audio_feat.extract_file(audio_path, &self.device)?; + let probs = if feats.dim(0)? <= self.cfg.chunk_max_frame { + let (probs, _) = self.vad_model.forward(&feats.unsqueeze(0)?, None)?; + probs + } else { + let mut chunk_probs = vec![]; + let chunks = split_tensor_with_size(&feats, self.cfg.chunk_max_frame, 0usize)?; + for chunk in chunks.iter() { + let (chunk_prob, _) = self.vad_model.forward(&chunk.unsqueeze(0)?, None)?; + chunk_probs.push(chunk_prob); + } + Tensor::cat(&chunk_probs, 1)? + }; + let probs = if self.model_name.to_lowercase().contains("aed") { + // only care speech + probs + .squeeze(0)? + .narrow(D::Minus1, 0, 1)? + .squeeze(D::Minus1)? + } else { + probs.squeeze(0)?.squeeze(D::Minus1)? + }; + let segments = self.vad_postprocessor.process(&probs, dur)?; + let res = VadResult { + dur, + timestamps: segments, + model_name: self.model_name.clone(), + mode: "speech".to_string(), + }; + Ok(res) + } + + pub fn reset(&mut self) { + self.caches = None; + } +} diff --git a/src/models/mod.rs b/src/models/mod.rs index dd4511b..8e3c440 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -24,6 +24,8 @@ pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; pub mod w2v_bert_2_0; +// pub mod sam3; +pub mod fire_red_vad; use crate::{ models::{ diff --git a/src/models/sam3/mod.rs b/src/models/sam3/mod.rs new file mode 100644 index 0000000..946b39c --- /dev/null +++ b/src/models/sam3/mod.rs @@ -0,0 +1 @@ +pub mod processor; \ No newline at end of file diff --git a/src/models/sam3/processor.rs b/src/models/sam3/processor.rs new file mode 100644 index 0000000..e69de29 diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs index 16265b0..4f6e3ba 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -534,7 +534,8 @@ impl VoxCPMModel { let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?; let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; - let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; + let mut audio = + load_audio_with_resample(&path, &self.device, Some(self.sample_rate), false)?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = @@ -574,7 +575,8 @@ impl VoxCPMModel { let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?; let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; - let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; + let mut audio = + load_audio_with_resample(&path, &self.device, Some(self.sample_rate), false)?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = @@ -841,8 +843,12 @@ impl VoxCPMModel { ) -> Result> { let text_token = self.tokenizer.encode(prompt_text)?; let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; - let mut audio = - load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?; + let mut audio = load_audio_with_resample( + &prompt_wav_path, + &self.device, + Some(self.sample_rate), + false, + )?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?; diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index eb9ac7a..169f4eb 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -462,7 +462,11 @@ pub fn get_audio_format_from_bytes(bytes: &[u8]) -> Result { } } -pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<(Tensor, usize)> { +pub fn load_audio_use_symphonia( + audio_vec: Vec, + is_i16: bool, + device: &Device, +) -> Result<(Tensor, usize)> { let extension = get_audio_format_from_bytes(&audio_vec)?; let content = Cursor::new(audio_vec); let mss = MediaSourceStream::new(Box::new(content), Default::default()); @@ -505,7 +509,14 @@ pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<( all_samples.push(Vec::new()); } let channel_data = buf.chan(channel); - all_samples[channel].extend_from_slice(channel_data); + if is_i16 { + // 将[-1.0, 1.0] => [-i16_max, i16_max] + let i16_data: Vec = + channel_data.iter().map(|&s| s * 32768.0).collect(); + all_samples[channel].extend_from_slice(&i16_data); + } else { + all_samples[channel].extend_from_slice(channel_data); + } } } AudioBufferRef::S16(buf) => { @@ -516,10 +527,17 @@ pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<( all_samples.push(Vec::new()); } let channel_data = buf.chan(channel); - let float_samples: Vec = channel_data - .iter() - .map(|&s| s as f32 / 32768.0) // 转换为[-1, 1] - .collect(); + let float_samples: Vec = if is_i16 { + channel_data + .iter() + .map(|&s| s as f32) // 转换为f32类型 + .collect() + } else { + channel_data + .iter() + .map(|&s| s as f32 / 32768.0) // 转换为[-1, 1] + .collect() + }; all_samples[channel].extend(float_samples); } } @@ -531,10 +549,17 @@ pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<( all_samples.push(Vec::new()); } let channel_data = buf.chan(channel); - let float_samples: Vec = channel_data - .iter() - .map(|&s| s.inner() as f32 / 8388608.0) // 转换为[-1, 1] - .collect(); + let float_samples: Vec = if is_i16 { + channel_data + .iter() + .map(|&s| s.inner() as f32 / 8388608.0 * 32768.0) // 转换为[-i16_max, i16_max] + .collect() + } else { + channel_data + .iter() + .map(|&s| s.inner() as f32 / 8388608.0) // 转换为[-1, 1] + .collect() + }; all_samples[channel].extend(float_samples); } } @@ -559,20 +584,16 @@ pub fn load_audio_use_symphonia(audio_vec: Vec, device: &Device) -> Result<( pub fn load_audio(path: &str, device: &Device) -> Result<(Tensor, usize)> { let audio_vec = get_audio_bytes_vec(path)?; - load_audio_use_symphonia(audio_vec, device) + load_audio_use_symphonia(audio_vec, false, device) } -pub fn load_audio_with_resample( - path: &str, +pub fn load_audio_with_resample_from_bytes( + audio_vec: Vec, device: &Device, target_sample_rate: Option, + is_i16: bool, ) -> Result { - // hound 只支持wav文件 - // let audio_path = get_audio_path(path)?; - // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; - - let audio_vec = get_audio_bytes_vec(path)?; - let (mut audio, sr) = load_audio_use_symphonia(audio_vec, device)?; + let (mut audio, sr) = load_audio_use_symphonia(audio_vec, is_i16, device)?; if let Some(target_sample_rate) = target_sample_rate && target_sample_rate != sr { @@ -581,6 +602,19 @@ pub fn load_audio_with_resample( Ok(audio) } +pub fn load_audio_with_resample( + path: &str, + device: &Device, + target_sample_rate: Option, + is_i16: bool, +) -> Result { + // hound 只支持wav文件 + // let audio_path = get_audio_path(path)?; + // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; + let audio_vec = get_audio_bytes_vec(path)?; + load_audio_with_resample_from_bytes(audio_vec, device, target_sample_rate, is_i16) +} + pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()> { let spec = hound::WavSpec { channels: 1, @@ -653,7 +687,7 @@ pub fn extract_audios( // 并行加载音频 audio_url_vec .par_iter() - .map(|url| load_audio_with_resample(url, device, target_sample_rate)) + .map(|url| load_audio_with_resample(url, device, target_sample_rate, false)) .collect() // #[cfg(not(feature = "ffmpeg"))] // { diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index fd86e7d..543c234 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -33,6 +33,28 @@ pub fn attn_masked_fill(on_true: &Tensor, mask: &Tensor, on_false: f32) -> Resul Ok(filled) } +pub fn get_mask_from_lengths(length: &Tensor) -> Result { + // length: [5u32, 4, 3, 6] + // mask: + // [[0, 0, 0, 0, 0, 1], + // [0, 0, 0, 0, 1, 1], + // [0, 0, 0, 1, 1, 1], + // [0, 0, 0, 0, 0, 0]] + let n = length.dim(0)?; + let t = length.max_all()?.to_scalar::()? as usize; + let mut mask = Tensor::zeros((n, t), DType::U32, length.device())?; + for i in 0..n { + let index = length.i(i)?.to_scalar::()? as usize; + let len = t - index; + if len == 0 { + continue; + } + let slice = Tensor::ones((1, len), DType::U32, length.device())?; + mask = mask.slice_assign(&[(i..i + 1), (index..t)], &slice)?; + } + Ok(mask) +} + pub fn prepare_mask(mask: &Tensor) -> Result { //(bs, seq_len) // [[1, 1, 1, 1, 0, 0]] @@ -584,3 +606,12 @@ pub fn repeat_interleave(t: &Tensor, repeats: usize, dim: usize) -> Result Result { + // probs shape: (m) + let m = probs.dim(0)?; + let thres_vec = vec![threshold; m]; + let thres_t = Tensor::new(thres_vec, probs.device())?; + let res = probs.ge(&thres_t)?; + Ok(res) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 91bbd05..ee0df5a 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -5,13 +5,15 @@ // use std::io::{Read, Seek}; // use std::{io::Cursor, time::Instant}; -use aha::{ - models::common::model_mapping::WhichModel, - // utils::{timestamp, timestamp_millis}, -}; +use aha::utils::tensor_utils::get_mask_from_lengths; // use aha::utils::tensor_utils::repeat_interleave; // use crate::params::chat::ChatCompletionParameters; -use anyhow::Result; +use anyhow::{Result}; +use candle_core::Tensor; +// use kaldi_native_fbank::{ +// FbankComputer, FbankOptions, +// window::{Window, extract_window}, +// }; // use byteorder::{LittleEndian, ReadBytesExt}; // use candle_core::Tensor; use modelscope::{DownloadOptions, ModelScope}; @@ -39,10 +41,76 @@ async fn download_test() -> Result<()> { #[test] fn messy_test() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda --test messy_test messy_test -r -- --nocapture - let model = WhichModel::LFM2_1_2B; - println!("model: {:?}, model_id: {}", model, model.as_string()); - let model_list = WhichModel::model_list(); - println!("model_list: {:#?}", model_list); + let device = aha::Device::Cpu; + let input = Tensor::new(&[5u32, 4, 3, 6], &device)?; + let mask = get_mask_from_lengths(&input)?; + println!("{}", mask); + // let audio_path = "file:///home/jhq/python_code/FireRedASR2S/assets/hello_zh.wav"; + // let device = aha::Device::Cpu; + // let wave = load_audio_with_resample(audio_path, &device, Some(16000), true)?; + // println!("len: {}", wave); + // let wave = wave.squeeze(0)?.to_vec1::()?; + // println!("wave len: {}", wave.len()); + // let mut opts = FbankOptions::default(); + // opts.frame_opts.dither = 0.0; + // opts.frame_opts.samp_freq = 16000.; + // opts.frame_opts.frame_length_ms = 25.; + // opts.frame_opts.frame_shift_ms = 10.; + // opts.frame_opts.snip_edges = true; + // opts.mel_opts.num_bins = 80; + // opts.mel_opts.debug_mel = false; + // opts.use_energy = false; + + // let mut comp = + // FbankComputer::new(opts.clone()).map_err(|e| anyhow!("fbank comput err: {e}"))?; + // let win = Window::new(&opts.frame_opts).unwrap(); + // let padded = opts.frame_opts.padded_window_size(); + + // let mut feats = vec![]; + // let mut window_buf = vec![0.0; padded]; + // for frame in 0..230 { + // let raw_log_energy = extract_window( + // 0, + // &wave, + // frame, + // &opts.frame_opts, + // Some(&win), + // &mut window_buf, + // ) + // .unwrap(); + // let mut feat = vec![0.0; comp.dim()]; + // comp.compute(raw_log_energy, 1.0, &mut window_buf, &mut feat); + // feats.push(feat); + // } + // let feats = Tensor::new(feats, &aha::Device::Cpu)?; + // println!("feats: {}", feats); + // println!("wave len: {}", wave.len()); + // let mut feats = vec![]; + // let frame_num = (wave.len() + padded - 1) / padded; + // for i in 0..frame_num { + // let mut window_buf = vec![0.0; padded]; + // let raw_log_energy = + // extract_window(0, &wave, i, &opts.frame_opts, Some(&win), &mut window_buf) + // .map_err(|_| anyhow!("extract_window err"))?; + // let mut feat = vec![0.0; comp.dim()]; + // comp.compute(raw_log_energy, 1.0, &mut window_buf, &mut feat); + // feats.push(feat); + // } + // let feats = Tensor::new(feats, &aha::Device::Cpu)?; + // println!("feats: {:?}", feats); + // let mut window_buf = vec![0.0; padded]; + // println!("window_buf len: {}", window_buf.len()); + // let raw_log_energy = extract_window(0, &wave, 0, &opts.frame_opts, Some(&win), &mut window_buf) + // .map_err(|_| anyhow!("extract_window err"))?; + + // let mut feat = vec![0.0; comp.dim()]; + // println!("feat len: {}", feat.len()); + // comp.compute(raw_log_energy, 1.0, &mut window_buf, &mut feat); + // println!("{feat:?}"); + // let model = WhichModel::LFM2_1_2B; + // println!("model: {:?}, model_id: {}", model, model.as_string()); + // let model_list = WhichModel::model_list(); + // println!("model_list: {:#?}", model_list); // println!("当前秒级时间戳: {}", timestamp()); // println!("当前毫秒级时间戳: {}", timestamp_millis()); diff --git a/tests/test_fire_red_vad.rs b/tests/test_fire_red_vad.rs new file mode 100644 index 0000000..c06d69c --- /dev/null +++ b/tests/test_fire_red_vad.rs @@ -0,0 +1,47 @@ +use aha::models::fire_red_vad::vad::FireRedVad; +use anyhow::Result; + +#[test] +fn aed() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_fire_red_vad aed -r -- --nocapture + let audio_path = "file:///home/jhq/python_code/FireRedASR2S/assets/hello_zh.wav"; + let device = aha::Device::Cpu; + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/xukaituo/FireRedVAD/AED/", save_dir); + + let vad = FireRedVad::init(&model_path, Some(&device), None)?; + let res = vad.detect_file(audio_path)?; + println!("vad res: {:?}", res); + Ok(()) +} + +#[test] +fn stream_vad() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_fire_red_vad stream_vad -r -- --nocapture + let audio_path = "file:///home/jhq/python_code/FireRedASR2S/assets/hello_zh.wav"; + let device = aha::Device::Cpu; + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/xukaituo/FireRedVAD/Stream-VAD/", save_dir); + + let vad = FireRedVad::init(&model_path, Some(&device), None)?; + let res = vad.detect_file(audio_path)?; + println!("vad res: {:?}", res); + Ok(()) +} + +#[test] +fn vad() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda --test test_fire_red_vad vad -r -- --nocapture + let audio_path = "file:///home/jhq/python_code/FireRedASR2S/assets/hello_zh.wav"; + let device = aha::Device::Cpu; + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/xukaituo/FireRedVAD/VAD/", save_dir); + + let vad = FireRedVad::init(&model_path, Some(&device), None)?; + let res = vad.detect_file(audio_path)?; + println!("vad res: {:?}", res); + Ok(()) +} diff --git a/tests/test_robo_brain.rs b/tests/test_robo_brain.rs index 7da265f..064fa56 100644 --- a/tests/test_robo_brain.rs +++ b/tests/test_robo_brain.rs @@ -1,11 +1,57 @@ use std::time::Instant; +use aha::models::qwen3vl::generate::Qwen3VLGenerateModel; use aha::models::{GenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel}; use aha::params::chat::ChatCompletionParameters; use anyhow::Result; #[test] -fn robo_brain_generate() -> Result<()> { +fn robo_brain2_5_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_robo_brain robo_brain2_5_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!("{}/BAAI/RoboBrain2.5-4B/", save_dir); + + let message = r#" + { + "model": "RoboBrain2.5-4B", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image", + "image_url": + { + "url": "http://images.cocodataset.org/val2017/000000039769.jpg" + } + }, + { + "type": "text", + "text": "What is shown in this image?" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut model = Qwen3VLGenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); + } + Ok(()) +} + +#[test] +fn robo_brain2_0_generate() -> Result<()> { // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda robo_brain_generate -r -- --nocapture let save_dir = diff --git a/tests/weight_test.rs b/tests/weight_test.rs index b81ff80..9e10bf3 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -351,3 +351,62 @@ fn voxcpm2_weight() -> Result<()> { Ok(()) } + +#[test] +fn sam3_weight() -> Result<()> { + // cargo test -F cuda --test weight_test sam3_weight -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/facebook/sam3/", save_dir); + let model_list = find_type_files(&model_path, "safetensors")?; + println!("model_list: {:?}", model_list); + let device = get_device(None); + // let mut dict_to_hashmap = HashMap::new(); + // let mut dtype = candle_core::DType::F32; + for m in model_list { + let weights = safetensors::load(m, &device)?; + for (key, tensor) in weights.iter() { + println!("=== {} === {:?}", key, tensor); + } + } + + Ok(()) +} + +#[test] +fn sam3_1_weight() -> Result<()> { + // cargo test -F cuda --test weight_test sam3_1_weight -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/facebook/sam3.1/", save_dir); + let model_list = find_type_files(&model_path, "pt")?; + println!("model_list: {:?}", model_list); + // let dev = get_device(None); + // let mut dict_to_hashmap = HashMap::new(); + // let mut dtype = candle_core::DType::F32; + for m in model_list { + let dict = read_all_with_key(m, None)?; + // dtype = dict[0].1.dtype(); + for (k, v) in dict { + println!("key: {}, tensor shape: {:?}", k, v); + // dict_to_hashmap.insert(k, v); + } + } + + Ok(()) +} + +#[test] +fn fire_red_vad_weight() -> Result<()> { + // cargo test -F cuda --test weight_test fire_red_vad_weight -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/xukaituo/FireRedVAD/VAD/model.safetensors", save_dir); + let device = get_device(None); + let weights = safetensors::load(model_path, &device)?; + for (key, tensor) in weights.iter() { + println!("=== {} === {:?}", key, tensor); + } + + Ok(()) +}