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
+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)
}