add FireRedVAD
This commit is contained in:
Generated
+29
-3
@@ -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"
|
||||
|
||||
@@ -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" }
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)?)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub mod config;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
pub mod vad;
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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::{
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
pub mod processor;
|
||||
@@ -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
@@ -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"))]
|
||||
// {
|
||||
|
||||
@@ -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
@@ -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());
|
||||
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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 =
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user