add FireRedVAD

This commit is contained in:
jhqxxx
2026-04-15 17:42:41 +08:00
parent 3e30360998
commit f901f4cbfe
22 changed files with 1244 additions and 37 deletions
Generated
+29 -3
View File
@@ -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"
+1
View File
@@ -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" }
+3
View File
@@ -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
+3
View File
@@ -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
+3
View File
@@ -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
+3
View File
@@ -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
+38
View File
@@ -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<Tensor> {
// 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<Tensor> {
Ok(t.log()?.affine(1.0 / 10.0_f64.ln(), 0.0)?)
}
+136
View File
@@ -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<Vec<f32>>,
}
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,
}
}
}
+4
View File
@@ -0,0 +1,4 @@
pub mod config;
pub mod model;
pub mod processor;
pub mod vad;
+291
View File
@@ -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<Conv1d>,
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<Self> {
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<Self> {
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<DFSMNBlock>,
dnns: Vec<Linear>, // 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<Self> {
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<Tensor>>,
) -> Result<(Tensor, Vec<Tensor>)> {
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<Self> {
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<Tensor>>,
) -> Result<(Tensor, Vec<Tensor>)> {
let (x, new_caches) = self.dfsmn.forward(feat, None, caches)?;
let logits = self.out.forward(&x)?;
let probs = sigmoid(&logits)?;
Ok((probs, new_caches))
}
}
+244
View File
@@ -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<Self> {
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::<f32>()?;
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<Tensor> {
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<Self> {
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<Tensor> {
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::<f32>()?
} else if wav_tensor.rank() == 2 {
wav_tensor.squeeze(0)?.to_vec1::<f32>()?
} 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<Self> {
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<Tensor> {
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<bool> {
// TODO 状态管理
let is_speech = probs >= self.prob_threshold;
Ok(is_speech)
}
pub fn process_thresh(&self, raw_probs: &Tensor) -> Result<Tensor> {
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<Vec<(f32, f32)>> {
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<Vec<(f32, f32)>> {
let mut segments = vec![];
let mut speech_start = -1;
let decisions = decisions.to_vec1::<u8>()?;
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<Tensor> {
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::<f32>()?;
mean_vec.push(mean);
}
let means = Tensor::new(mean_vec, probs.device())?;
moothed = moothed.slice_assign(&[(0..mean_len)], &means)?;
Ok(moothed)
}
}
}
+161
View File
@@ -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<Tensor>,
pub kaldi_audio: Option<Tensor>,
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<Vec<Tensor>>,
frame_length_sample: usize,
}
impl FireRedVad {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
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<u8>) -> Result<Option<VadFrameResult>> {
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::<u32>()?;
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<VadResult> {
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;
}
}
+2
View File
@@ -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::{
+1
View File
@@ -0,0 +1 @@
pub mod processor;
View File
+10 -4
View File
@@ -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<HashMap<String, Tensor>> {
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)?;
+54 -20
View File
@@ -462,7 +462,11 @@ pub fn get_audio_format_from_bytes(bytes: &[u8]) -> Result<String> {
}
}
pub fn load_audio_use_symphonia(audio_vec: Vec<u8>, device: &Device) -> Result<(Tensor, usize)> {
pub fn load_audio_use_symphonia(
audio_vec: Vec<u8>,
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<u8>, 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<f32> =
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<u8>, device: &Device) -> Result<(
all_samples.push(Vec::new());
}
let channel_data = buf.chan(channel);
let float_samples: Vec<f32> = channel_data
.iter()
.map(|&s| s as f32 / 32768.0) // 转换为[-1, 1]
.collect();
let float_samples: Vec<f32> = 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<u8>, device: &Device) -> Result<(
all_samples.push(Vec::new());
}
let channel_data = buf.chan(channel);
let float_samples: Vec<f32> = channel_data
.iter()
.map(|&s| s.inner() as f32 / 8388608.0) // 转换为[-1, 1]
.collect();
let float_samples: Vec<f32> = 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<u8>, 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<u8>,
device: &Device,
target_sample_rate: Option<usize>,
is_i16: bool,
) -> Result<Tensor> {
// 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<usize>,
is_i16: bool,
) -> Result<Tensor> {
// 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"))]
// {
+31
View File
@@ -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<Tensor> {
// 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::<u32>()? 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::<u32>()? 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<Tensor> {
//(bs, seq_len)
// [[1, 1, 1, 1, 0, 0]]
@@ -584,3 +606,12 @@ pub fn repeat_interleave(t: &Tensor, repeats: usize, dim: usize) -> Result<Tenso
let t = t.index_select(&indices_tensor, dim)?;
Ok(t)
}
pub fn apply_threshold(probs: &Tensor, threshold: f32) -> Result<Tensor> {
// 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)
}
+77 -9
View File
@@ -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::<f32>()?;
// 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());
+47
View File
@@ -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(())
}
+47 -1
View File
@@ -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 =
+59
View File
@@ -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(())
}