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