From d9b803d27ec47382c81659c670d584d5a55e8e78 Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Thu, 15 Jan 2026 21:57:12 +0800 Subject: [PATCH] add qwen3 and fun-asr-nano --- Cargo.lock | 22 +- Cargo.toml | 25 +- README.md | 11 +- download_and_run.sh | 16 +- src/main.rs | 2 + src/models/common/mod.rs | 94 +++- src/models/deepseek_ocr/model.rs | 3 + src/models/fun_asr_nano/config.rs | 84 ++++ src/models/fun_asr_nano/generate.rs | 207 +++++++++ src/models/fun_asr_nano/mod.rs | 4 + src/models/fun_asr_nano/model.rs | 646 +++++++++++++++++++++++++++ src/models/fun_asr_nano/processor.rs | 110 +++++ src/models/glm_asr_nano/generate.rs | 2 +- src/models/glm_asr_nano/processor.rs | 50 +-- src/models/minicpm4/model.rs | 3 + src/models/mod.rs | 25 +- src/models/qwen3/config.rs | 44 ++ src/models/qwen3/generate.rs | 172 +++++++ src/models/qwen3/mod.rs | 3 + src/models/qwen3/model.rs | 277 ++++++++++++ src/models/qwen3vl/config.rs | 40 +- src/models/qwen3vl/generate.rs | 11 +- src/models/qwen3vl/model.rs | 187 +------- src/models/voxcpm/generate.rs | 28 +- src/models/voxcpm/minicpm4.rs | 3 + src/models/voxcpm/model.rs | 8 +- src/position_embed/mod.rs | 1 + src/position_embed/sinusoidal_pe.rs | 59 +++ src/utils/audio_utils.rs | 414 ++++++++++++++++- src/utils/img_utils.rs | 50 ++- src/utils/tensor_utils.rs | 39 ++ tests/test_fun_asr_nano.rs | 97 ++++ tests/test_glm_asr_nano.rs | 2 +- tests/test_paddleocr_vl.rs | 4 +- tests/test_qwen3.rs | 83 ++++ tests/test_qwen3vl.rs | 2 +- tests/test_voxcpm.rs | 4 +- tests/test_voxcpm1_5.rs | 8 +- tests/weight_test.rs | 44 ++ 39 files changed, 2577 insertions(+), 307 deletions(-) create mode 100644 src/models/fun_asr_nano/config.rs create mode 100644 src/models/fun_asr_nano/generate.rs create mode 100644 src/models/fun_asr_nano/mod.rs create mode 100644 src/models/fun_asr_nano/model.rs create mode 100644 src/models/fun_asr_nano/processor.rs create mode 100644 src/models/qwen3/config.rs create mode 100644 src/models/qwen3/generate.rs create mode 100644 src/models/qwen3/mod.rs create mode 100644 src/models/qwen3/model.rs create mode 100644 src/position_embed/sinusoidal_pe.rs create mode 100644 tests/test_fun_asr_nano.rs create mode 100644 tests/test_qwen3.rs diff --git a/Cargo.lock b/Cargo.lock index 943e39b..d09013d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -19,7 +19,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" [[package]] name = "aha" -version = "0.1.7" +version = "0.1.8" dependencies = [ "aha_openai_dive", "anyhow", @@ -43,6 +43,7 @@ dependencies = [ "rocket", "serde", "serde_json", + "serde_yaml", "symphonia", "tokenizers", "tokio", @@ -3718,6 +3719,19 @@ dependencies = [ "serde", ] +[[package]] +name = "serde_yaml" +version = "0.9.34+deprecated" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6a8b1a1a2ebf674015cc02edccce75287f1a0130d394307b36743c2f5d504b47" +dependencies = [ + "indexmap", + "itoa", + "ryu", + "serde", + "unsafe-libyaml", +] + [[package]] name = "sharded-slab" version = "0.1.7" @@ -4633,6 +4647,12 @@ version = "0.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "323402cff2dd658f39ca17c789b502021b3f18707c91cdf22e3838e1b4023817" +[[package]] +name = "unsafe-libyaml" +version = "0.2.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "673aac59facbab8a9007c7f6108d11f63b603f7cabff99fabf650fea5c32b861" + [[package]] name = "untrusted" version = "0.9.0" diff --git a/Cargo.toml b/Cargo.toml index f72347f..2f5bdbd 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,15 +1,15 @@ [package] name = "aha" -version = "0.1.7" +version = "0.1.8" edition = "2024" repository = "https://github.com/jhqxxx/aha" license = "Apache-2.0" -description = "aha model inference library, now supports Qwen2.5VL, MiniCPM4, VoxCPM, Qwen3VL, DeepSeek-OCR, Hunyuan-OCR, PaddleOCR-VL, VoxCPM1.5, RMBG2.0, GLM-ASR-Nano-2512" +description = "aha model inference library, now supports Qwen2.5VL, MiniCPM4, VoxCPM, Qwen3VL, DeepSeek-OCR, Hunyuan-OCR, PaddleOCR-VL, VoxCPM1.5, RMBG2.0, GLM-ASR-Nano-2512, Fun-ASR-Nano-2512, Qwen3" [dependencies] -candle-core = { version = "0.9.1"} -candle-nn = { version = "0.9.1"} -candle-transformers = { version = "0.9.1"} +candle-core = { version = "0.9.1" } +candle-nn = { version = "0.9.1" } +candle-transformers = { version = "0.9.1" } candle-flash-attn = { version = "0.9.1", optional = true } serde = "1.0.226" serde_json = "1.0.145" @@ -21,9 +21,9 @@ base64 = "0.22.1" num = "0.4.3" minijinja = "2.12.0" tokenizers = "0.22.1" -aha_openai_dive = { version = "1.4", features = ["stream"]} -uuid = { version = "1.18.1", features = ["v4"]} -chrono = "0.4.42" +aha_openai_dive = { version = "1.4", features = ["stream"] } +uuid = { version = "1.18.1", features = ["v4"] } +chrono = "0.4" rocket = { version = "0.5.1", features = ["serde_json", "json"] } tokio = "1.47.1" hound = "3.5.1" @@ -36,12 +36,13 @@ rayon = "1.10" # audioadapter-buffers = "2.0.0" realfft = "3.5.0" symphonia = { version = "0.5.5", features = ["mp3", "wav"] } +serde_yaml = "0.9.34" [features] -flash-attn=["candle-flash-attn"] -cuda=["candle-nn/cuda", "candle-core/cuda", "candle-transformers/cuda"] -metal=["candle-nn/metal", "candle-core/metal", "candle-transformers/metal"] -ffmpeg=["ffmpeg-next"] +flash-attn = ["candle-flash-attn"] +cuda = ["candle-nn/cuda", "candle-core/cuda", "candle-transformers/cuda"] +metal = ["candle-nn/metal", "candle-core/metal", "candle-transformers/metal"] +ffmpeg = ["ffmpeg-next"] [lints.clippy] needless_range_loop = "allow" diff --git a/README.md b/README.md index 68dfd13..d3a959d 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,10 @@ - 模型:[RMBG2.0](https://huggingface.co/briaai/RMBG-2.0) 开源协议:[Attribution-NonCommercial 4.0 International](https://creativecommons.org/licenses/by-nc/4.0/deed.en) * GLM-ASR-Nano-2512 - 智谱AI语音识别模型 - 模型:[GLM-ASR-Nano-2512](https://huggingface.co/zai-org/GLM-ASR-Nano-2512) 开源协议:[MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) +* Fun-ASR-Nano-2512 - 通义百聆语音识别模型 + - 模型:[Fun-ASR-Nano-2512](https://huggingface.co/FunAudioLLM/Fun-ASR-Nano-2512) 开源协议未标明 +* [Qwen3](https://huggingface.co/collections/Qwen/qwen3) - 通义千问 Qwen3系列语言模型 + - 模型:[Qwen3-0.6B](https://huggingface.co/Qwen/Qwen3-0.6B) 开源协议: [Apache license 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) ## 计划支持 我们持续扩展支持的模型列表,欢迎贡献! @@ -106,6 +110,7 @@ cargo run -F cuda -r -- [参数] * minicpm4-0.5b:OpenBMB/MiniCPM4-0.5B 模型 * qwen2.5vl-3b:Qwen/Qwen2.5-VL-3B-Instruct 模型 * qwen2.5vl-7b:Qwen/Qwen2.5-VL-7B-Instruct 模型 + * qwen3-0.6b: Qwen/Qwen3-0.6B 模型 * qwen3vl-2b:Qwen/Qwen3-VL-2B-Instruct 模型 * qwen3vl-4b:Qwen/Qwen3-VL-4B-Instruct 模型 * qwen3vl-8b:Qwen/Qwen3-VL-8B-Instruct 模型 @@ -117,6 +122,7 @@ cargo run -F cuda -r -- [参数] * voxcpm: OpenBMB/VoxCPM-0.5B 模型 * voxcpm1.5: OpenBMB/VoxCPM1.5 模型 * glm-asr-nano-2512: ZhipuAI/GLM-ASR-Nano-2512 模型 + * fun-asr-nano-2512: FunAudioLLM/Fun-ASR-Nano-2512 模型 * 示例:--model deepseek-ocr 或 -m qwen3vl-2b 3. 权重路径 @@ -153,7 +159,7 @@ cargo run -F cuda -r -- [参数] 1. 对话接口 - **端点**: `POST /chat/completions` - **功能**: 多模态对话和文本生成 -- **支持模型**: Qwen2.5VL,Qwen3VL,DeepSeekOCR, GLM-ASR-Nano-2512 等 +- **支持模型**: Qwen2.5VL, Qwen3, Qwen3VL, DeepSeekOCR, GLM-ASR-Nano-2512, Fun-ASR-Nano-2512 等 - **请求格式**: OpenAI Chat Completion 格式 - **响应格式**: OpenAI Chat Completion 格式 - **流式支持**: 支持 @@ -289,6 +295,9 @@ cargo test -F cuda voxcpm_generate -r -- --nocapture 2. 提交新的 Issue,包含详细描述和复现步骤 ## 更新日志 +### v0.1.8 +* 支持Fun-ASR-Nano-2512, Qwen3 模型 + ### v0.1.7 * 支持GLM-ASR-Nano-2512 模型 diff --git a/download_and_run.sh b/download_and_run.sh index 867cb87..684aab1 100755 --- a/download_and_run.sh +++ b/download_and_run.sh @@ -16,7 +16,8 @@ show_help() { echo "Available models:" echo " minicpm4-0.5b" echo " qwen2.5vl-3b" - echo " qwen2.5vl-7b" + echo " qwen2.5vl-7b" + echo " qwen3-0.6b" echo " qwen3vl-2b" echo " qwen3vl-4b" echo " qwen3vl-8b" @@ -28,6 +29,7 @@ show_help() { echo " voxcpm" echo " voxcpm1.5" echo " glm-asr-nano-2512" + echo " fun-asr-nano-2512" echo "" exit 1 } @@ -51,6 +53,9 @@ case $MODEL_ALIAS in "qwen2.5vl-7b") MODEL_ID="Qwen/Qwen2.5-VL-7B-Instruct" ;; + "qwen3-0.6b") + MODEL_ID="Qwen/Qwen3-0.6B" + ;; "qwen3vl-2b") MODEL_ID="Qwen/Qwen3-VL-2B-Instruct" ;; @@ -73,17 +78,20 @@ case $MODEL_ALIAS in MODEL_ID="PaddlePaddle/PaddleOCR-VL" ;; "RMBG2.0") - MODEL_ID="AI-ModelScope/RMBG-2.0" + MODEL_ID="briaai/RMBG-2.0" ;; "voxcpm") - MODEL_ID="OpenBMB/VoxCPM-0.5B" + MODEL_ID="openbmb/VoxCPM-0.5B" ;; "voxcpm1.5") - MODEL_ID="OpenBMB/VoxCPM1.5" + MODEL_ID="openbmb/VoxCPM1.5" ;; "glm-asr-nano-2512") MODEL_ID="zai-org/GLM-ASR-Nano-2512" ;; + "fun-asr-nano-2512") + MODEL_ID="FunAudioLLM/Fun-ASR-Nano-2512" + ;; *) echo "Error: Unknown model alias '$MODEL_ALIAS'" show_help diff --git a/src/main.rs b/src/main.rs index 2126c5d..be1118e 100644 --- a/src/main.rs +++ b/src/main.rs @@ -74,6 +74,7 @@ async fn main() -> anyhow::Result<()> { WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B", WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct", WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct", + WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B", WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct", WhichModel::Qwen3vl4B => "Qwen/Qwen3-VL-4B-Instruct", WhichModel::Qwen3vl8B => "Qwen/Qwen3-VL-8B-Instruct", @@ -85,6 +86,7 @@ async fn main() -> anyhow::Result<()> { WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B", WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5", WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512", + WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512", }; let model_path = match &args.weight_path { Some(path) => path.clone(), diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index f72672e..737ce7e 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -1,4 +1,4 @@ -use anyhow::Result; +use anyhow::{Result, anyhow}; use candle_core::{D, Tensor}; use candle_nn::{ Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, Embedding, @@ -6,6 +6,7 @@ use candle_nn::{ conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear, linear_no_bias, rms_norm, }; +use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator}; use crate::{ position_embed::rope::{RoPE, apply_rotary_pos_emb}, @@ -126,6 +127,9 @@ impl NaiveAttention { num_key_value_heads: usize, head_dim: Option, bias: bool, + q_proj_pp_name: Option<&str>, + k_proj_pp_name: Option<&str>, + v_proj_pp_name: Option<&str>, o_proj_pp_name: Option<&str>, ) -> Result { let num_kv_groups = num_attention_heads / num_key_value_heads; @@ -133,12 +137,27 @@ impl NaiveAttention { None => hidden_size / num_attention_heads, Some(dim) => dim, }; + let q_proj_pp_name = q_proj_pp_name.unwrap_or("q_proj"); + let k_proj_pp_name = k_proj_pp_name.unwrap_or("k_proj"); + let v_proj_pp_name = v_proj_pp_name.unwrap_or("v_proj"); let o_proj_pp_name = o_proj_pp_name.unwrap_or("o_proj"); let (q_proj, k_proj, v_proj, o_proj) = if bias { ( - linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?, - linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?, - linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?, + linear( + hidden_size, + num_attention_heads * head_dim, + vb.pp(q_proj_pp_name), + )?, + linear( + hidden_size, + num_key_value_heads * head_dim, + vb.pp(k_proj_pp_name), + )?, + linear( + hidden_size, + num_key_value_heads * head_dim, + vb.pp(v_proj_pp_name), + )?, linear( num_attention_heads * head_dim, hidden_size, @@ -147,9 +166,21 @@ impl NaiveAttention { ) } else { ( - linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?, - linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?, - linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?, + linear_no_bias( + hidden_size, + num_attention_heads * head_dim, + vb.pp(q_proj_pp_name), + )?, + linear_no_bias( + hidden_size, + num_key_value_heads * head_dim, + vb.pp(k_proj_pp_name), + )?, + linear_no_bias( + hidden_size, + num_key_value_heads * head_dim, + vb.pp(v_proj_pp_name), + )?, linear_no_bias( num_attention_heads * head_dim, hidden_size, @@ -305,6 +336,9 @@ impl NaiveAttnTwoLinearMLPBlock { num_key_value_heads, head_dim, attn_bias, + None, + None, + None, o_proj_pp_name, )?; let mlp = TwoLinearMLP::new( @@ -386,6 +420,9 @@ impl NaiveAttnGateUpDownMLPBlock { num_key_value_heads, head_dim, attn_bias, + None, + None, + None, o_proj_pp_name, )?; let mlp = GateUpDownMLP::new( @@ -811,3 +848,46 @@ impl LlamaForCausalLM { self.model.clear_kv_cache(); } } + +pub fn conv1d_group_parallel(xs: &Tensor, conv1d: &Conv1d) -> Result { + let groups = conv1d.config().groups; + let xs = if groups == 1 { + xs.conv1d_with_algo( + conv1d.weight(), + conv1d.config().padding, + conv1d.config().stride, + conv1d.config().dilation, + groups, + conv1d.config().cudnn_fwd_algo, + )? + } else { + let blocks = xs.chunk(groups, 1)?; + let kernel = conv1d.weight().chunk(groups, 0)?; + let blocks = blocks + // .iter() + .par_iter() + .zip(&kernel) + .map(|(block, kernel)| { + block + .conv1d_with_algo( + kernel, + conv1d.config().padding, + conv1d.config().stride, + conv1d.config().dilation, + 1, + conv1d.config().cudnn_fwd_algo, + ) + .map_err(|e| anyhow!(format!("tensor conv1d_with_algo error:{}", e))) + }) + .collect::>>()?; + Tensor::cat(&blocks, 1)? + }; + match conv1d.bias() { + None => Ok(xs), + Some(bias) => { + let b = bias.dims1()?; + let bias = bias.reshape((1, b, 1))?; + Ok(xs.broadcast_add(&bias)?) + } + } +} diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs index 26764d1..4352a95 100644 --- a/src/models/deepseek_ocr/model.rs +++ b/src/models/deepseek_ocr/model.rs @@ -986,6 +986,9 @@ impl DeepseekV2DecoderLayer { None, false, None, + None, + None, + None, )?; let mlp = if layer_id >= config.first_k_dense_replace && layer_id.is_multiple_of(config.moe_layer_freq) diff --git a/src/models/fun_asr_nano/config.rs b/src/models/fun_asr_nano/config.rs new file mode 100644 index 0000000..2a40225 --- /dev/null +++ b/src/models/fun_asr_nano/config.rs @@ -0,0 +1,84 @@ +use serde::Deserialize; + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct FunASRNanoConfig { + pub audio_encoder_conf: AudioEncoderConf, + pub llm_conf: LlmConf, + pub audio_adaptor_conf: AudioAdaptorConf, + pub detach_ctc_decoder: bool, + pub ctc_decoder_conf: CtcDecoderConf, + pub ctc_weight: f64, + pub ctc_conf: CtcConf, + pub frontend_conf: FrontendConf, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct AudioEncoderConf { + pub output_size: usize, + pub attention_heads: usize, + pub linear_units: usize, + pub num_blocks: usize, + pub tp_blocks: usize, + pub dropout_rate: f64, + pub positional_dropout_rate: f64, + pub attention_dropout_rate: f64, + pub input_layer: String, + pub pos_enc_class: String, + pub normalize_before: bool, + pub kernel_size: usize, + pub sanm_shfit: usize, + pub selfattention_layer_type: String, + pub freeze: bool, + pub freeze_layer_num: i32, + pub feat_permute: bool, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct LlmConf { + pub hub: String, + pub freeze: bool, + pub llm_dtype: String, + pub init_param_path: String, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct AudioAdaptorConf { + pub downsample_rate: usize, + pub use_low_frame_rate: bool, + pub ffn_dim: usize, + pub llm_dim: usize, + pub encoder_dim: usize, + pub n_layer: usize, + pub freeze: bool, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct CtcDecoderConf { + pub downsample_rate: u32, + pub ffn_dim: u32, + pub llm_dim: u32, + pub encoder_dim: u32, + pub n_layer: u32, + pub freeze: bool, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct CtcConf { + pub dropout_rate: f64, + pub ctc_type: String, + pub reduce: bool, + pub ignore_nan_grad: bool, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct FrontendConf { + pub fs: usize, + pub window: String, + pub n_mels: usize, + pub frame_length: f32, + pub frame_shift: f32, + pub lfr_m: usize, + pub lfr_n: usize, + #[serde(skip_serializing_if = "Option::is_none")] + pub cmvn_file: Option, +} diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs new file mode 100644 index 0000000..0b8864f --- /dev/null +++ b/src/models/fun_asr_nano/generate.rs @@ -0,0 +1,207 @@ +use std::collections::HashMap; + +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, Tensor, pickle::read_all_with_key}; +use candle_nn::VarBuilder; +use rocket::async_stream::stream; +use rocket::futures::Stream; + +use crate::{ + models::{ + GenerateModel, + fun_asr_nano::{ + config::FunASRNanoConfig, model::FunAsrNanoModel, processor::FunAsrNanoProcessor, + }, + qwen3::config::{Qwen3Config, Qwen3GenerationConfig}, + }, + tokenizer::TokenizerModel, + utils::{ + build_completion_chunk_response, build_completion_response, find_type_files, get_device, + get_dtype, get_logit_processor, + }, +}; + +pub struct FunAsrNanoGenerateModel { + tokenizer: TokenizerModel, + processor: FunAsrNanoProcessor, + fun_asr_nano: FunAsrNanoModel, + device: Device, + dtype: DType, + eos_token_id1: u32, + eos_token_id2: u32, + generation_config: Qwen3GenerationConfig, + model_name: String, +} + +impl FunAsrNanoGenerateModel { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let llm_config_path = path.to_string() + "/Qwen3-0.6B"; + let tokenizer = TokenizerModel::init(&llm_config_path)?; + let generation_config_path = llm_config_path.clone() + "/generation_config.json"; + let generation_config: Qwen3GenerationConfig = + serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let config_path = llm_config_path + "/config.json"; + let llm_cfg: Qwen3Config = serde_json::from_slice(&std::fs::read(config_path)?)?; + let device = get_device(device); + let config_path = path.to_string() + "/config.yaml"; + let cfg: FunASRNanoConfig = serde_yaml::from_slice(&std::fs::read(config_path)?)?; + let cfg_dtype = cfg.llm_conf.llm_dtype.as_str(); + let dtype = get_dtype(dtype, cfg_dtype); + let processor = FunAsrNanoProcessor::new(&cfg.frontend_conf, &device)?; + let model_list = find_type_files(path, "pt")?; + let mut dict_to_hashmap = HashMap::new(); + for m in model_list { + let dict = read_all_with_key(m, Some("state_dict"))?; + for (k, v) in dict { + dict_to_hashmap.insert(k, v); + } + } + let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device); + let fun_asr_nano = FunAsrNanoModel::new(vb, &cfg, &llm_cfg)?; + Ok(Self { + tokenizer, + processor, + fun_asr_nano, + device, + dtype, + eos_token_id1: generation_config.eos_token_id[0] as u32, + eos_token_id2: generation_config.eos_token_id[1] as u32, + generation_config, + model_name: "fun-asr-nano".to_string(), + }) + } +} + +impl GenerateModel for FunAsrNanoGenerateModel { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let mut logit_processor = + get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); + let (speech, fbank_mask, mut input_ids) = + self.processor.process_info(&mes, &self.tokenizer)?; + let mut speech = Some(speech.to_dtype(self.dtype)?); + let mut fbank_mask = Some(&fbank_mask); + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let mut generate = Vec::new(); + let sample_len = mes.max_tokens.unwrap_or(1024); + for _ in 0..sample_len { + let logits = self.fun_asr_nano.forward( + &input_ids, + speech.as_ref(), + fbank_mask, + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + generate.push(next_token); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + speech = None; + fbank_mask = None; + } + let num_token = generate.len() as u32; + let res = self.tokenizer.token_decode(generate)?; + self.fun_asr_nano.clear_kv_cache(); + let response = build_completion_response(res, &self.model_name, Some(num_token)); + Ok(response) + } + + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let mut logit_processor = + get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); + let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let sample_len = mes.max_tokens.unwrap_or(1024); + let stream = stream! { + let mut error_tokens = Vec::new(); + let mut speech = Some(speech.to_dtype(self.dtype)?); + let mut fbank_mask = Some(&fbank_mask); + let mut input_ids = input_ids; + for _ in 0..sample_len { + let logits = self.fun_asr_nano.forward( + &input_ids, + speech.as_ref(), + fbank_mask, + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + let mut decode_ids = Vec::new(); + if !error_tokens.is_empty() { + decode_ids.extend_from_slice(&error_tokens); + } + decode_ids.push(next_token); + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; + if decoded_token.contains("�") { + error_tokens.push(next_token); + if error_tokens.len() > 3 { + error_tokens.clear(); + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + speech = None; + fbank_mask = None; + continue; + } + error_tokens.clear(); + let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); + yield Ok(chunk); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + speech = None; + fbank_mask = None; + } + self.fun_asr_nano.clear_kv_cache(); + }; + Ok(Box::new(Box::pin(stream))) + } +} diff --git a/src/models/fun_asr_nano/mod.rs b/src/models/fun_asr_nano/mod.rs new file mode 100644 index 0000000..8b1baf7 --- /dev/null +++ b/src/models/fun_asr_nano/mod.rs @@ -0,0 +1,4 @@ +pub mod config; +pub mod generate; +pub mod model; +pub mod processor; diff --git a/src/models/fun_asr_nano/model.rs b/src/models/fun_asr_nano/model.rs new file mode 100644 index 0000000..d3f971f --- /dev/null +++ b/src/models/fun_asr_nano/model.rs @@ -0,0 +1,646 @@ +use anyhow::Result; +use candle_core::{D, IndexOp, Tensor}; +use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::softmax_last_dim}; + +use crate::{ + models::{ + common::{ + NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + }, + fun_asr_nano::config::FunASRNanoConfig, + qwen3::{config::Qwen3Config, model::Qwen3Model}, + }, + position_embed::sinusoidal_pe::SinusoidalPositionEncoderCat, + utils::tensor_utils::{get_equal_mask, mask_filled, masked_scatter_dim0}, +}; + +pub struct MultiHeadedAttentionSANM { + head_dim: usize, + n_head: usize, + linear_out: Linear, + linear_q_k_v: Linear, + fsmn_block: Conv1d, + left_padding: usize, + right_padding: usize, + scaling: f64, +} + +impl MultiHeadedAttentionSANM { + pub fn new( + vb: VarBuilder, + n_head: usize, + in_dim: usize, + hidden_dim: usize, + kernel_size: usize, + sanm_shfit: usize, + ) -> Result { + let head_dim = hidden_dim / n_head; + let linear_out = linear(hidden_dim, hidden_dim, vb.pp("linear_out"))?; + let linear_q_k_v = linear(in_dim, hidden_dim * 3, vb.pp("linear_q_k_v"))?; + let fsmn_block = get_conv1d( + vb.pp("fsmn_block"), + hidden_dim, + hidden_dim, + kernel_size, + 0, + 1, + 1, + hidden_dim, + false, + )?; + let mut left_padding = (kernel_size - 1) / 2; + if sanm_shfit > 0 { + left_padding += sanm_shfit; + } + let right_padding = kernel_size - 1 - left_padding; + let scaling = (head_dim as f64).powf(-0.5); + Ok(Self { + head_dim, + n_head, + linear_out, + linear_q_k_v, + fsmn_block, + left_padding, + right_padding, + scaling, + }) + } + + pub fn forward_fsmn( + &self, + inputs: &Tensor, + mask: Option<&Tensor>, + mask_shfit_chunk: Option<&Tensor>, + ) -> Result { + let mut inputs = inputs.clone(); + let mask = if let Some(mask) = mask { + let mut mask = mask.unsqueeze(D::Minus1)?.unsqueeze(0)?; + if let Some(mask_shfit_chunk) = mask_shfit_chunk { + mask = mask.broadcast_mul(mask_shfit_chunk)?; + } + inputs = inputs.broadcast_mul(&mask)?; + Some(mask) + } else { + None + }; + let xs = inputs.transpose(1, 2)?; + let xs = xs.pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)?; + let xs = self.fsmn_block.forward(&xs)?; + let xs = xs.transpose(1, 2)?; + let mut xs = xs.add(&inputs)?; + if let Some(mask) = mask { + xs = xs.broadcast_mul(&mask)?; + } + Ok(xs) + } + pub fn forward_qkv(&self, xs: &Tensor) -> Result<(Tensor, Tensor, Tensor, Tensor)> { + let (b, t, _) = xs.dims3()?; + let q_k_v = self + .linear_q_k_v + .forward(xs)? + .reshape((b, t, 3, self.n_head, ()))? + .permute((2, 0, 3, 1, 4))? + .contiguous()?; + let q_h = q_k_v.i(0)?.contiguous()?; + let k_h = q_k_v.i(1)?.contiguous()?; + let v_h = q_k_v.i(2)?.contiguous()?; + let v = v_h.transpose(1, 2)?.reshape((b, t, ()))?; + Ok((q_h, k_h, v_h, v)) + } + + pub fn forward_attention( + &self, + values: &Tensor, + scores: &Tensor, + mask: Option<&Tensor>, + mask_att_chunk_encoder: Option<&Tensor>, + ) -> Result { + let bs = scores.dim(0)?; + let attn = if let Some(mask) = mask { + let mask = if let Some(mask_att_chunk_encoder) = mask_att_chunk_encoder { + mask.mul(mask_att_chunk_encoder)? + } else { + mask.clone() + }; + // mask: rank = 2 + let mask = get_equal_mask(&mask, 0)?; + let scores = mask_filled(scores, &mask, f32::NEG_INFINITY)?; + let attn = softmax_last_dim(&scores)?; + mask_filled(&attn, &mask, 0.0)? + } else { + softmax_last_dim(scores)? + }; + let xs = attn.matmul(values)?; + let xs = + xs.transpose(1, 2)? + .contiguous()? + .reshape((bs, (), self.n_head * self.head_dim))?; + let xs = self.linear_out.forward(&xs)?; + Ok(xs) + } + + pub fn forward_simple(&self, xs: &Tensor) -> Result { + let (b, t, _) = xs.dims3()?; + let q_k_v = self.linear_q_k_v.forward(xs)?; + let dim = self.head_dim * self.n_head; + let q_h = q_k_v + .narrow(D::Minus1, 0, dim)? + .reshape((b, t, self.n_head, ()))? + .permute((0, 2, 1, 3))?; + let k_h = q_k_v + .narrow(D::Minus1, dim, dim)? + .reshape((b, t, self.n_head, ()))? + .permute((0, 2, 1, 3))?; + let v = q_k_v.narrow(D::Minus1, dim * 2, dim)?; + let v_h = v.reshape((b, t, self.n_head, ()))?.permute((0, 2, 1, 3))?; + let fsmn_memory = v.transpose(1, 2)?; + let fsmn_memory = fsmn_memory + .pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)? + .contiguous()?; + let fsmn_memory = self.fsmn_block.forward(&fsmn_memory)?; + // let fsmn_memory = conv1d_group_parallel(&fsmn_memory, &self.fsmn_block)?; + + let fsmn_memory = fsmn_memory.transpose(1, 2)?; + let fsmn_memory = fsmn_memory.add(&v)?; + let att_outs = eager_attention_forward(&q_h, &k_h, &v_h, None, None, self.scaling)?; + let att_outs = att_outs.reshape((b, t, ()))?; + let att_outs = self.linear_out.forward(&att_outs)?; + let att_outs = att_outs.add(&fsmn_memory)?; + Ok(att_outs) + } + + pub fn forward( + &self, + xs: &Tensor, + mask: Option<&Tensor>, + mask_shfit_chunk: Option<&Tensor>, + mask_att_chunk_encoder: Option<&Tensor>, + ) -> Result { + let (q_h, k_h, v_h, v) = self.forward_qkv(xs)?; + let fsmn_memory = self.forward_fsmn(&v, mask, mask_shfit_chunk)?; + let q_h = q_h.affine(self.scaling, 0.0)?; + let scores = q_h.matmul(&k_h.transpose(D::Minus2, D::Minus1)?)?; + let attn_outs = self.forward_attention(&v_h, &scores, mask, mask_att_chunk_encoder)?; + let att_outs = attn_outs.add(&fsmn_memory)?; + Ok(att_outs) + } +} + +pub struct EncoderLayerSANM { + self_attn: MultiHeadedAttentionSANM, + feed_forward: TwoLinearMLP, + norm1: LayerNorm, + norm2: LayerNorm, + concat_linear: Option, + normalize_before: bool, + in_dim: usize, + hidden_dim: usize, +} + +impl EncoderLayerSANM { + pub fn new( + vb: VarBuilder, + in_dim: usize, + hidden_dim: usize, + n_head: usize, + kernel_size: usize, + sanm_shfit: usize, + hidden_units: usize, + normalize_before: bool, + concat_after: bool, + ) -> Result { + let self_attn = MultiHeadedAttentionSANM::new( + vb.pp("self_attn"), + n_head, + in_dim, + hidden_dim, + kernel_size, + sanm_shfit, + )?; + let feed_forward = TwoLinearMLP::new( + vb.pp("feed_forward"), + hidden_dim, + hidden_units, + hidden_dim, + candle_nn::Activation::Relu, + true, + "w_1", + "w_2", + )?; + let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, in_dim)?; + let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, hidden_dim)?; + let concat_linear = if concat_after { + let lin = linear(hidden_dim * 2, hidden_dim, vb.pp("concat_linear"))?; + Some(lin) + } else { + None + }; + Ok(Self { + self_attn, + feed_forward, + norm1, + norm2, + concat_linear, + normalize_before, + in_dim, + hidden_dim, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + mask: Option<&Tensor>, + mask_shfit_chunk: Option<&Tensor>, + mask_att_chunk_encoder: Option<&Tensor>, + ) -> Result { + let stoch_layer_coeff = 1.0f64; + let residual = xs.clone(); + let mut xs = if self.normalize_before { + self.norm1.forward(xs)? + } else { + xs.clone() + }; + if self.concat_linear.is_some() { + let attn = + self.self_attn + .forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?; + let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?; + if self.in_dim == self.hidden_dim { + let x_concat = self + .concat_linear + .as_ref() + .unwrap() + .forward(&x_concat)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&x_concat)?; + } else { + xs = self + .concat_linear + .as_ref() + .unwrap() + .forward(&x_concat)? + .affine(stoch_layer_coeff, 0.0)?; + } + } else if self.in_dim == self.hidden_dim { + let attn = self + .self_attn + .forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&attn)?; + } else { + xs = self + .self_attn + .forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)? + .affine(stoch_layer_coeff, 0.0)?; + } + + if !self.normalize_before { + xs = self.norm1.forward(&xs)?; + } + let residual = xs.clone(); + if self.normalize_before { + xs = self.norm2.forward(&xs)?; + } + xs = self + .feed_forward + .forward(&xs)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&xs)?; + if !self.normalize_before { + xs = self.norm2.forward(&xs)?; + } + Ok(xs) + } + + pub fn forward_simple(&self, xs: &Tensor) -> Result { + let residual = xs.clone(); + let mut xs = self.norm1.forward(xs)?; + if self.in_dim == self.hidden_dim { + let attn = self.self_attn.forward_simple(&xs)?; + xs = residual.add(&attn)?; + } else { + xs = self.self_attn.forward_simple(&xs)?; + } + + let residual = xs.clone(); + let xs = self.norm2.forward(&xs)?; + + let xs = self.feed_forward.forward(&xs)?; + let xs = residual.add(&xs)?; + Ok(xs) + } +} + +pub struct SenseVoiceEncoderSmall { + embed: SinusoidalPositionEncoderCat, + encoders0: EncoderLayerSANM, + encoders: Vec, + tp_encoders: Vec, + after_norm: LayerNorm, + tp_norm: LayerNorm, + scaling: f64, +} + +impl SenseVoiceEncoderSmall { + pub fn new( + vb: VarBuilder, + input_size: usize, + output_size: usize, + attention_heads: usize, + linear_units: usize, + num_blocks: usize, + tp_blocks: usize, + normalize_before: bool, + kernel_size: usize, + sanm_shfit: usize, + ) -> Result { + let embed = SinusoidalPositionEncoderCat::new(Some(input_size), true, vb.device())?; + + let encoders0 = EncoderLayerSANM::new( + vb.pp("encoders0.0"), + input_size, + output_size, + attention_heads, + kernel_size, + sanm_shfit, + linear_units, + normalize_before, + false, + )?; + let mut encoders = vec![]; + let vb_encoders = vb.pp("encoders"); + for i in 0..(num_blocks - 1) { + let encoder_i = EncoderLayerSANM::new( + vb_encoders.pp(i), + output_size, + output_size, + attention_heads, + kernel_size, + sanm_shfit, + linear_units, + normalize_before, + false, + )?; + encoders.push(encoder_i); + } + let vb_tp_encoders = vb.pp("tp_encoders"); + let mut tp_encoders = vec![]; + for i in 0..tp_blocks { + let tp_blocks_i = EncoderLayerSANM::new( + vb_tp_encoders.pp(i), + output_size, + output_size, + attention_heads, + kernel_size, + sanm_shfit, + linear_units, + normalize_before, + false, + )?; + tp_encoders.push(tp_blocks_i); + } + let after_norm = get_layer_norm(vb.pp("after_norm"), 1e-5, output_size)?; + let tp_norm = get_layer_norm(vb.pp("tp_norm"), 1e-5, output_size)?; + let scaling = (output_size as f64).powf(0.5); + Ok(Self { + embed, + encoders0, + encoders, + tp_encoders, + after_norm, + tp_norm, + scaling, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = xs.affine(self.scaling, 0.0)?; + let xs = self.embed.forward(&xs, 0)?; + let mut xs = self.encoders0.forward_simple(&xs)?; + for encoder_layer in &self.encoders { + xs = encoder_layer.forward_simple(&xs)?; + } + xs = self.after_norm.forward(&xs)?; + for tp_layer in &self.tp_encoders { + xs = tp_layer.forward_simple(&xs)?; + } + xs = self.tp_norm.forward(&xs)?; + Ok(xs) + } +} + +pub struct AdaptorEncoderLayer { + self_attn: NaiveAttention, + feed_forward: TwoLinearMLP, + norm1: LayerNorm, + norm2: LayerNorm, + concat_linear: Option, + normalize_before: bool, +} + +impl AdaptorEncoderLayer { + pub fn new( + vb: VarBuilder, + llm_dim: usize, + n_head: usize, + normalize_before: bool, + concat_after: bool, + ) -> Result { + let self_attn = NaiveAttention::new( + vb.pp("self_attn"), + llm_dim, + n_head, + n_head, + None, + true, + Some("linear_q"), + Some("linear_k"), + Some("linear_v"), + Some("linear_out"), + )?; + let feed_forward = TwoLinearMLP::new( + vb.pp("feed_forward"), + llm_dim, + llm_dim / 4, + llm_dim, + candle_nn::Activation::Relu, + true, + "w_1", + "w_2", + )?; + let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, llm_dim)?; + let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, llm_dim)?; + let concat_linear = if concat_after { + let lin = linear(llm_dim * 2, llm_dim, vb.pp("concat_linear"))?; + Some(lin) + } else { + None + }; + Ok(Self { + self_attn, + feed_forward, + norm1, + norm2, + concat_linear, + normalize_before, + }) + } + + pub fn forward(&self, xs: &Tensor, mask: Option<&Tensor>) -> Result { + let stoch_layer_coeff = 1.0f64; + let residual = xs.clone(); + let mut xs = if self.normalize_before { + self.norm1.forward(xs)? + } else { + xs.clone() + }; + if self.concat_linear.is_some() { + let attn = self.self_attn.forward(&xs, None, None, mask, false)?; + let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?; + let x_concat = self + .concat_linear + .as_ref() + .unwrap() + .forward(&x_concat)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&x_concat)?; + } else { + let attn = self + .self_attn + .forward(&xs, None, None, mask, false)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&attn)?; + } + if !self.normalize_before { + xs = self.norm1.forward(&xs)?; + } + let residual = xs.clone(); + if self.normalize_before { + xs = self.norm2.forward(&xs)?; + } + xs = self + .feed_forward + .forward(&xs)? + .affine(stoch_layer_coeff, 0.0)?; + xs = residual.add(&xs)?; + if !self.normalize_before { + xs = self.norm2.forward(&xs)?; + } + Ok(xs) + } +} + +pub struct AudioAdaptor { + k: usize, + linear1: Linear, + linear2: Linear, + blocks: Vec, +} + +impl AudioAdaptor { + pub fn new( + vb: VarBuilder, + downsample_rate: usize, + encoder_dim: usize, + llm_dim: usize, + ffn_dim: usize, + n_layer: usize, + attention_heads: usize, + ) -> Result { + let linear1 = linear(encoder_dim * downsample_rate, ffn_dim, vb.pp("linear1"))?; + let linear2 = linear(ffn_dim, llm_dim, vb.pp("linear2"))?; + let mut blocks = vec![]; + let vb_blocks = vb.pp("blocks"); + for i in 0..n_layer { + let layer = + AdaptorEncoderLayer::new(vb_blocks.pp(i), llm_dim, attention_heads, true, false)?; + blocks.push(layer); + } + Ok(Self { + k: downsample_rate, + linear1, + linear2, + blocks, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let (bs, seq_len, dim) = xs.dims3()?; + let chunk_num = (seq_len - 1) / self.k + 1; + let pad_num = chunk_num * self.k - seq_len; + let xs = xs.pad_with_zeros(1, 0, pad_num)?; + let xs = xs.contiguous()?.reshape((bs, chunk_num, dim * self.k))?; + let xs = self.linear1.forward(&xs)?.relu()?; + let mut xs = self.linear2.forward(&xs)?; + for block in &self.blocks { + xs = block.forward(&xs, None)?; + } + Ok(xs) + } +} + +pub struct FunAsrNanoModel { + audio_encoder: SenseVoiceEncoderSmall, + audio_adaptor: AudioAdaptor, + llm: Qwen3Model, +} +impl FunAsrNanoModel { + pub fn new(vb: VarBuilder, config: &FunASRNanoConfig, llm_cfg: &Qwen3Config) -> Result { + let input_size = config.frontend_conf.lfr_m * config.frontend_conf.n_mels; + let audio_encoder = SenseVoiceEncoderSmall::new( + vb.pp("audio_encoder"), + input_size, + config.audio_encoder_conf.output_size, + config.audio_encoder_conf.attention_heads, + config.audio_encoder_conf.linear_units, + config.audio_encoder_conf.num_blocks, + config.audio_encoder_conf.tp_blocks, + config.audio_encoder_conf.normalize_before, + config.audio_encoder_conf.kernel_size, + config.audio_encoder_conf.sanm_shfit, + )?; + let audio_adaptor = AudioAdaptor::new( + vb.pp("audio_adaptor"), + config.audio_adaptor_conf.downsample_rate, + config.audio_adaptor_conf.encoder_dim, + config.audio_adaptor_conf.llm_dim, + config.audio_adaptor_conf.ffn_dim, + config.audio_adaptor_conf.n_layer, + 8, + )?; + let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?; + Ok(Self { + audio_encoder, + audio_adaptor, + llm, + }) + } + + pub fn forward( + &mut self, + input_ids: &Tensor, + speech: Option<&Tensor>, + fbank_mask: Option<&Tensor>, + seqlen_offset: usize, + ) -> Result { + let mut inputs_embeds = self.llm.embedding_token_id(input_ids)?; + if let Some(speech) = speech + && let Some(fbank_mask) = fbank_mask + { + let speech = self.audio_encoder.forward(speech)?; + let encoder_out = self.audio_adaptor.forward(&speech)?; + let speech_token_len = fbank_mask.sum_all()?.to_scalar::()?; + let audio_embed = encoder_out + .squeeze(0)? + .narrow(0, 0, speech_token_len as usize)?; + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embed, fbank_mask)?; + } + let logits = self + .llm + .forward(None, Some(&inputs_embeds), seqlen_offset)?; + Ok(logits) + } + + pub fn clear_kv_cache(&mut self) { + self.llm.clear_kv_cache(); + } +} diff --git a/src/models/fun_asr_nano/processor.rs b/src/models/fun_asr_nano/processor.rs new file mode 100644 index 0000000..c44b46d --- /dev/null +++ b/src/models/fun_asr_nano/processor.rs @@ -0,0 +1,110 @@ +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use candle_core::{D, Device, Tensor}; + +use crate::{ + models::fun_asr_nano::config::FrontendConf, + tokenizer::TokenizerModel, + utils::{ + audio_utils::{ + apply_lfr, extract_audios, get_waveform_and_window_properties, kaldi_fbank, + kaldi_get_mel_banks, + }, + extract_user_text, + }, +}; + +pub struct FunAsrNanoProcessor { + fronted_conf: FrontendConf, + device: Device, + prompt_prefix: String, + prompt_suffix: String, + window_shift: usize, + window_size: usize, + padded_window_size: usize, + mel_energies: Tensor, +} + +impl FunAsrNanoProcessor { + pub fn new(fronted_conf: &FrontendConf, device: &Device) -> Result { + let (window_shift, window_size, padded_window_size) = get_waveform_and_window_properties( + fronted_conf.fs, + fronted_conf.frame_shift, + fronted_conf.frame_length, + true, + )?; + let (mel_energies, _) = kaldi_get_mel_banks( + fronted_conf.n_mels, + padded_window_size, + fronted_conf.fs as f32, + 20.0, + 0.0, + device, + )?; + let mel_energies = mel_energies.pad_with_zeros(D::Minus1, 0, 1)?.t()?; + Ok(Self { + fronted_conf: fronted_conf.clone(), + device: device.clone(), + prompt_prefix: + "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n" + .to_string(), + prompt_suffix: "<|im_end|>\n<|im_start|>assistant\n".to_string(), + window_shift, + window_size, + padded_window_size, + mel_energies, + }) + } + + pub fn extract_fbank(&self, audio: &Tensor) -> Result<(Tensor, usize)> { + let waveform = audio.affine(32768.0, 0.0)?; + let mut mat = kaldi_fbank( + &waveform, + &self.mel_energies, + self.window_shift, + self.window_size, + self.padded_window_size, + 1.0, + // 0.0, + // "hamming", + // self.fronted_conf.fs, + // true, + )?; + mat = mat.squeeze(0)?; + if self.fronted_conf.lfr_m != 1 || self.fronted_conf.lfr_n != 1 { + mat = apply_lfr(&mat, self.fronted_conf.lfr_m, self.fronted_conf.lfr_n)?; + } + let feat_length = mat.dim(0)?; + let mat = mat.unsqueeze(0)?; + Ok((mat, feat_length)) + } + + pub fn process_info( + &self, + mes: &ChatCompletionParameters, + tokenizer: &TokenizerModel, + ) -> Result<(Tensor, Tensor, Tensor)> { + let user_text = extract_user_text(mes)?; + let sub_prompt = self.prompt_prefix.clone() + &user_text; + let mut source_ids = vec![]; + let mut fbank_mask = vec![]; + let sub_token = tokenizer.text_encode_vec(sub_prompt, true)?; + source_ids.extend_from_slice(&sub_token); + fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]); + let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?; + let audio = &audio_tensors[0]; + let (speech, speech_lengths) = self.extract_fbank(audio)?; + let olens = 1 + (speech_lengths - 3 + 2) / 2; + let olens = 1 + (olens - 3 + 2) / 2; + let fake_token_len = (olens - 1) / 2 + 1; + source_ids.extend_from_slice(&vec![0u32; fake_token_len]); + fbank_mask.extend_from_slice(&vec![1u32; fake_token_len]); + let sub_token = tokenizer.text_encode_vec(self.prompt_suffix.clone(), true)?; + source_ids.extend_from_slice(&sub_token); + fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]); + let input_ids = Tensor::from_slice(&source_ids, (1, source_ids.len()), &self.device)?; + let fbank_mask = Tensor::from_slice(&fbank_mask, (1, fbank_mask.len()), &self.device)?; + + Ok((speech, fbank_mask, input_ids)) + } +} diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index 5674531..b96ace4 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -70,7 +70,7 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { Some(s) => s as u64, }; let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); - let render_text = self.chat_template.apply_chat_template(&mes)?; + let render_text: String = self.chat_template.apply_chat_template(&mes)?; let (input_features, audio_token_lengths, replace_text) = self.processor.process_info(&mes, &render_text)?; let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs index 043a756..5ba02da 100644 --- a/src/models/glm_asr_nano/processor.rs +++ b/src/models/glm_asr_nano/processor.rs @@ -3,12 +3,13 @@ use std::f32; use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::Result; use candle_core::{D, DType, Device, IndexOp, Tensor}; -use rayon::iter::{IntoParallelRefIterator, ParallelIterator}; use crate::{ models::glm_asr_nano::config::GlmAsrNanoProcessorConfig, utils::{ - audio_utils::{create_hann_window, extract_audios, mel_filter_bank, stft_audio}, + audio_utils::{ + apply_stft, create_hann_window, extract_audios, extract_frames, mel_filter_bank, + }, tensor_utils::{pad_reflect_last_dim, split_tensor}, }, }; @@ -81,48 +82,19 @@ impl GlmAsrNanoProcessor { }) } - /// 提取音频帧 - pub fn extract_frames(&self, waveform: &Tensor, n_frames: usize) -> Result { - let mut frames = Vec::with_capacity(n_frames); - - for i in 0..n_frames { - let start = i * self.hop_length; - let frame = waveform.narrow(D::Minus1, start, self.n_fft)?; - frames.push(frame); - } - - let result = Tensor::cat(&frames, D::Minus1)?; - let bs = result.dim(0)?; - let reshaped = result.reshape((bs, n_frames, self.n_fft))?; - Ok(reshaped) - } - pub fn extract_fbank_features(&self, waveform: &Tensor) -> Result { let pad = self.n_fft / 2; let waveform = pad_reflect_last_dim(waveform, (pad, pad))?; - let (batch_size, samples) = waveform.dims2()?; + let (_, samples) = waveform.dims2()?; // 计算输出维度 let n_frames = (samples - self.n_fft) / self.hop_length + 1; // (bs, n_frames, n_fft) - let frames = self.extract_frames(&waveform, n_frames)?; + let frames = extract_frames(&waveform, self.n_fft, self.hop_length)?; // 应用汉明窗口 let result = frames.broadcast_mul(&self.window)?; // 傅立叶变换 - let mut wave_fft = vec![]; - for bs in 0..batch_size { - let wave_i = result.i(bs)?; - let wave_i_vec = wave_i.to_vec2::()?; - let wave_i_fft_vec: Result>> = wave_i_vec - .par_iter() - .map(|frame_wave| stft_audio(self.n_fft, frame_wave)) - .collect(); - let wave_i_fft_vec = wave_i_fft_vec?; - - let wave_i_fft = Tensor::new(wave_i_fft_vec, &self.device)?.unsqueeze(0)?; - wave_fft.push(wave_i_fft); - } - let magnitudes = Tensor::cat(&wave_fft, 0)?.transpose(D::Minus1, D::Minus2)?; + let magnitudes = apply_stft(&result)?.transpose(D::Minus1, D::Minus2)?; let magnitudes = magnitudes.narrow(D::Minus1, 0, n_frames - 1)?; let mel_spec = self.mel_filters.broadcast_matmul(&magnitudes)?; let mel_spec = mel_spec.clamp(1e-10f32, f32::INFINITY)?; @@ -141,12 +113,18 @@ impl GlmAsrNanoProcessor { let audio_len = audio.dim(0)?; let pad_num = self.n_samples - audio_len; - let audio_pad = audio.pad_with_zeros(0, 0, pad_num)?; + let audio_pad = if pad_num > 0 { + audio.pad_with_zeros(0, 0, pad_num)? + } else { + audio + }; // (n_samples) -> (1, n_samples) let audio_pad = audio_pad.unsqueeze(0)?; pad_audio.push(audio_pad); let mut mask = vec![1u32; audio_len]; - mask.extend_from_slice(&vec![0u32; pad_num]); + if pad_num > 0 { + mask.extend_from_slice(&vec![0u32; pad_num]); + } input_features_mask.push(mask); } let input_features = Tensor::cat(&pad_audio, 0)?; diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index 4a9ac61..535f964 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -111,6 +111,9 @@ impl MiniCPMDecoderLayer { None, false, None, + None, + None, + None, )?; let mlp = GateUpDownMLP::new( vb.pp("mlp"), diff --git a/src/models/mod.rs b/src/models/mod.rs index 9e4affb..4165669 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -1,10 +1,12 @@ pub mod common; pub mod deepseek_ocr; +pub mod fun_asr_nano; pub mod glm_asr_nano; pub mod hunyuan_ocr; pub mod minicpm4; pub mod paddleocr_vl; pub mod qwen2_5vl; +pub mod qwen3; pub mod qwen3vl; pub mod rmbg2_0; pub mod voxcpm; @@ -17,11 +19,12 @@ use rocket::futures::Stream; use crate::models::{ deepseek_ocr::generate::DeepseekOCRGenerateModel, + fun_asr_nano::generate::FunAsrNanoGenerateModel, glm_asr_nano::generate::GlmAsrNanoGenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, - qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model, - voxcpm::generate::VoxCPMGenerate, + qwen3::generate::Qwen3GenerateModel, qwen3vl::generate::Qwen3VLGenerateModel, + rmbg2_0::generate::RMBG2_0Model, voxcpm::generate::VoxCPMGenerate, }; #[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)] @@ -32,6 +35,8 @@ pub enum WhichModel { Qwen2_5vl3B, #[value(name = "qwen2.5vl-7b")] Qwen2_5vl7B, + #[value(name = "qwen3-0.6b")] + Qwen3_0_6B, #[value(name = "qwen3vl-2b")] Qwen3vl2B, #[value(name = "qwen3vl-4b")] @@ -54,6 +59,8 @@ pub enum WhichModel { VoxCPM1_5, #[value(name = "glm-asr-nano-2512")] GlmASRNano2512, + #[value(name = "fun-asr-nano-2512")] + FunASRNano2512, } pub trait GenerateModel { @@ -74,6 +81,7 @@ pub trait GenerateModel { pub enum ModelInstance<'a> { MiniCPM4(MiniCPMGenerateModel<'a>), Qwen2_5VL(Qwen2_5VLGenerateModel<'a>), + Qwen3(Qwen3GenerateModel<'a>), Qwen3VL(Qwen3VLGenerateModel<'a>), DeepSeekOCR(DeepseekOCRGenerateModel), HunyuanOCR(HunyuanOCRGenerateModel<'a>), @@ -81,6 +89,7 @@ pub enum ModelInstance<'a> { RMBG2_0(Box), VoxCPM(Box), GlmASRNano(GlmAsrNanoGenerateModel<'a>), + FunASRNano(FunAsrNanoGenerateModel), } impl<'a> GenerateModel for ModelInstance<'a> { @@ -88,6 +97,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { match self { ModelInstance::MiniCPM4(model) => model.generate(mes), ModelInstance::Qwen2_5VL(model) => model.generate(mes), + ModelInstance::Qwen3(model) => model.generate(mes), ModelInstance::Qwen3VL(model) => model.generate(mes), ModelInstance::DeepSeekOCR(model) => model.generate(mes), ModelInstance::HunyuanOCR(model) => model.generate(mes), @@ -95,6 +105,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::RMBG2_0(model) => model.generate(mes), ModelInstance::VoxCPM(model) => model.generate(mes), ModelInstance::GlmASRNano(model) => model.generate(mes), + ModelInstance::FunASRNano(model) => model.generate(mes), } } @@ -112,6 +123,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { match self { ModelInstance::MiniCPM4(model) => model.generate_stream(mes), ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes), + ModelInstance::Qwen3(model) => model.generate_stream(mes), ModelInstance::Qwen3VL(model) => model.generate_stream(mes), ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes), ModelInstance::HunyuanOCR(model) => model.generate_stream(mes), @@ -119,6 +131,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::RMBG2_0(model) => model.generate_stream(mes), ModelInstance::VoxCPM(model) => model.generate_stream(mes), ModelInstance::GlmASRNano(model) => model.generate_stream(mes), + ModelInstance::FunASRNano(model) => model.generate_stream(mes), } } } @@ -137,6 +150,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result { + let model = Qwen3GenerateModel::init(path, None, None)?; + ModelInstance::Qwen3(model) + } WhichModel::Qwen3vl2B => { let model = Qwen3VLGenerateModel::init(path, None, None)?; ModelInstance::Qwen3VL(model) @@ -181,6 +198,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result { + let model = FunAsrNanoGenerateModel::init(path, None, None)?; + ModelInstance::FunASRNano(model) + } }; Ok(model) } diff --git a/src/models/qwen3/config.rs b/src/models/qwen3/config.rs new file mode 100644 index 0000000..0cecd88 --- /dev/null +++ b/src/models/qwen3/config.rs @@ -0,0 +1,44 @@ +use candle_nn::Activation; +use serde::Deserialize; + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct Qwen3Config { + pub attention_bias: bool, + pub attention_dropout: f64, + pub bos_token_id: u32, + pub eos_token_id: u32, + pub head_dim: usize, + pub hidden_act: Activation, + pub hidden_size: usize, + pub initializer_range: f64, + pub intermediate_size: usize, + pub max_position_embeddings: usize, + pub max_window_layers: usize, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + pub rms_norm_eps: f64, + pub rope_theta: f32, + pub tie_word_embeddings: bool, + pub torch_dtype: String, + pub use_cache: bool, + pub use_sliding_window: bool, + pub vocab_size: usize, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3GenerationConfig { + pub bos_token_id: usize, + pub pad_token_id: usize, + pub do_sample: bool, + pub eos_token_id: Vec, + pub top_p: f32, + pub top_k: usize, + pub temperature: f32, + #[serde(default = "default_repetition_penalty")] + pub repetition_penalty: f32, +} + +fn default_repetition_penalty() -> f32 { + 1.0 +} diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs new file mode 100644 index 0000000..fa035bd --- /dev/null +++ b/src/models/qwen3/generate.rs @@ -0,0 +1,172 @@ +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, Tensor}; +use candle_nn::VarBuilder; +use rocket::async_stream::stream; +use rocket::futures::Stream; + +use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig}; +use crate::models::qwen3::model::Qwen3Model; +// use crate::models::GenerateStream; +use crate::utils::{ + build_completion_chunk_response, build_completion_response, find_type_files, get_device, + get_dtype, get_logit_processor, +}; +use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; + +pub struct Qwen3GenerateModel<'a> { + chat_template: ChatTemplate<'a>, + tokenizer: TokenizerModel, + qwen3: Qwen3Model, + device: Device, + eos_token_id1: u32, + eos_token_id2: u32, + generation_config: Qwen3GenerationConfig, + model_name: String, +} + +impl<'a> Qwen3GenerateModel<'a> { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let chat_template = ChatTemplate::init(path)?; + let tokenizer = TokenizerModel::init(path)?; + let config_path = path.to_string() + "/config.json"; + let cfg: Qwen3Config = serde_json::from_slice(&std::fs::read(config_path)?)?; + let device = &get_device(device); + let cfg_dtype = cfg.torch_dtype.as_str(); + let dtype = get_dtype(dtype, cfg_dtype); + let model_list = find_type_files(path, "safetensors")?; + let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; + let qwen3 = Qwen3Model::new(&cfg, vb)?; + let generation_config_path = path.to_string() + "/generation_config.json"; + let generation_config: Qwen3GenerationConfig = + serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + + Ok(Qwen3GenerateModel { + chat_template, + tokenizer, + qwen3, + device: device.clone(), + eos_token_id1: generation_config.eos_token_id[0] as u32, + eos_token_id2: generation_config.eos_token_id[1] as u32, + generation_config, + model_name: "qwen3".to_string(), + }) + } +} + +impl<'a> GenerateModel for Qwen3GenerateModel<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let mut logit_processor = + get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); + let mes_render = self.chat_template.apply_chat_template(&mes)?; + let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let mut generate = Vec::new(); + let sample_len = mes.max_tokens.unwrap_or(2048); + for _ in 0..sample_len { + let logits = self.qwen3.forward(Some(&input_ids), None, seqlen_offset)?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + generate.push(next_token); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + } + let num_token = generate.len() as u32; + let res = self.tokenizer.token_decode(generate)?; + self.qwen3.clear_kv_cache(); + let response = build_completion_response(res, &self.model_name, Some(num_token)); + Ok(response) + } + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let mut logit_processor = + get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); + let mes_render = self.chat_template.apply_chat_template(&mes)?; + let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let sample_len = mes.max_tokens.unwrap_or(512); + let stream = stream! { + let mut error_tokens = Vec::new(); + for _ in 0..sample_len { + let logits = self.qwen3.forward( + Some(&input_ids), + None, + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + let mut decode_ids = Vec::new(); + if !error_tokens.is_empty(){ + decode_ids.extend_from_slice(&error_tokens); + } + decode_ids.push(next_token); + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; + if decoded_token.contains("�") { + error_tokens.push(next_token); + if error_tokens.len() > 3 { + error_tokens.clear(); + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + continue; + } + error_tokens.clear(); + let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); + yield Ok(chunk); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + + } + self.qwen3.clear_kv_cache(); + }; + Ok(Box::new(Box::pin(stream))) + } +} diff --git a/src/models/qwen3/mod.rs b/src/models/qwen3/mod.rs new file mode 100644 index 0000000..7fca417 --- /dev/null +++ b/src/models/qwen3/mod.rs @@ -0,0 +1,3 @@ +pub mod config; +pub mod generate; +pub mod model; diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs new file mode 100644 index 0000000..8292c2a --- /dev/null +++ b/src/models/qwen3/model.rs @@ -0,0 +1,277 @@ +use anyhow::Result; +use candle_core::Tensor; +use candle_nn::{ + Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear, linear_no_bias, rms_norm, +}; + +use crate::{ + models::{ + common::{GateUpDownMLP, eager_attention_forward}, + qwen3::config::Qwen3Config, + }, + position_embed::rope::{RoPE, apply_rotary_pos_emb}, + utils::tensor_utils::prepare_causal_attention_mask, +}; + +pub struct Qwen3Attention { + q_proj: Linear, + k_proj: Linear, + v_proj: Linear, + o_proj: Linear, + q_norm: RmsNorm, + k_norm: RmsNorm, + num_attention_heads: usize, + num_key_value_heads: usize, + num_kv_groups: usize, + head_dim: usize, + scaling: f64, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl Qwen3Attention { + pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { + let hidden_size = config.hidden_size; + let num_attention_heads = config.num_attention_heads; + let head_dim = config.head_dim; + let num_key_value_heads = config.num_key_value_heads; + let num_kv_groups = num_attention_heads / num_key_value_heads; + let scaling = 1f64 / f64::sqrt(head_dim as f64); + let (q_proj, k_proj, v_proj, o_proj) = if config.attention_bias { + let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?; + let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?; + let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?; + let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?; + (q_proj, k_proj, v_proj, o_proj) + } else { + let q_proj = + linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?; + let k_proj = + linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?; + let v_proj = + linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?; + let o_proj = + linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?; + (q_proj, k_proj, v_proj, o_proj) + }; + let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?; + let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?; + Ok(Self { + q_proj, + k_proj, + v_proj, + o_proj, + q_norm, + k_norm, + num_attention_heads, + num_key_value_heads, + num_kv_groups, + head_dim, + scaling, + kv_cache: None, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let (b_sz, q_len, _) = xs.dims3()?; + let query_states = self.q_proj.forward(xs)?.reshape(( + b_sz, + q_len, + self.num_attention_heads, + self.head_dim, + ))?; + let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?; + let key_states = self.k_proj.forward(xs)?.reshape(( + b_sz, + q_len, + self.num_key_value_heads, + self.head_dim, + ))?; + let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?; + let value_states = self.v_proj.forward(xs)?; + let value_states = value_states + .reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))? + .transpose(1, 2)?; + let (query_states, key_states) = + apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?; + let (key_states, value_states) = match &self.kv_cache { + None => (key_states, value_states), + Some((prev_k, prev_v)) => { + let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; + let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; + (key_states, value_states) + } + }; + self.kv_cache = Some((key_states.clone(), value_states.clone())); + let attn_output = eager_attention_forward( + &query_states, + &key_states, + &value_states, + Some(self.num_kv_groups), + attention_mask, + self.scaling, + )?; + let attn_output = + attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?; + let attn_output = attn_output.apply(&self.o_proj)?; + Ok(attn_output) + } + + pub fn clear_kv_cache(&mut self) { + self.kv_cache = None + } +} + +pub struct Qwen3DecoderLayer { + self_attn: Qwen3Attention, + mlp: GateUpDownMLP, + input_layernorm: RmsNorm, + post_attention_layernorm: RmsNorm, +} + +impl Qwen3DecoderLayer { + pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { + let self_attn = Qwen3Attention::new(config, vb.pp("self_attn"))?; + let mlp = GateUpDownMLP::new( + vb.pp("mlp"), + config.hidden_size, + config.intermediate_size, + config.hidden_act, + false, + )?; + let input_layernorm = rms_norm( + config.hidden_size, + config.rms_norm_eps, + vb.pp("input_layernorm"), + )?; + let post_attention_layernorm = rms_norm( + config.hidden_size, + config.rms_norm_eps, + vb.pp("post_attention_layernorm"), + )?; + Ok(Self { + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let xs = self.input_layernorm.forward(xs)?; + let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?; + let xs = residual.add(&xs)?; + let residual = xs.clone(); + let xs = self.post_attention_layernorm.forward(&xs)?; + let xs = self.mlp.forward(&xs)?; + let xs = residual.add(&xs)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + self.self_attn.clear_kv_cache(); + } +} + +pub struct Qwen3Model { + embed_tokens: Embedding, + layers: Vec, + norm: RmsNorm, + rotary_emb: RoPE, + lm_head: Linear, +} + +impl Qwen3Model { + pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { + let vb = vb.pp("model"); + let vocab_size = config.vocab_size; + let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?; + let mut layers = vec![]; + let vb_l = vb.pp("layers"); + for layer_idx in 0..config.num_hidden_layers { + let layer = Qwen3DecoderLayer::new(config, vb_l.pp(layer_idx))?; + layers.push(layer) + } + let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?; + let head_dim = config.head_dim; + let rotary_emb = RoPE::new(head_dim, config.rope_theta, vb.device())?; + let lm_head = if config.tie_word_embeddings { + Linear::new(embed_tokens.embeddings().clone(), None) + } else { + linear_no_bias(config.hidden_size, config.vocab_size, vb.pp("lm_head"))? + }; + Ok(Self { + embed_tokens, + layers, + norm, + rotary_emb, + lm_head, + }) + } + pub fn forward( + &mut self, + input_ids: Option<&Tensor>, + inputs_embeds: Option<&Tensor>, + seqlen_offset: usize, + ) -> Result { + if input_ids.is_none() && inputs_embeds.is_none() { + return Err(anyhow::anyhow!( + "You must specify exactly one of input_ids or inputs_embeds" + )); + } + let inputs_embeds = if let Some(inputs_embeds) = inputs_embeds { + inputs_embeds.clone() + } else { + let input_ids = input_ids.unwrap(); + self.embedding_token_id(input_ids)? + }; + let (bs, seq_len, _) = inputs_embeds.dims3()?; + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + bs, + seq_len, + 0, + inputs_embeds.device(), + )?) + } + }; + + let (cos, sin) = self + .rotary_emb + .forward(seqlen_offset, seq_len, inputs_embeds.device())?; + + let mut hidden_states = inputs_embeds; + for decode_layer in &mut self.layers { + hidden_states = + decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?; + } + hidden_states = self.norm.forward(&hidden_states)?; + let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&hidden_state)?; + Ok(logits) + } + pub fn embedding_token_id(&self, input_ids: &Tensor) -> Result { + Ok(self.embed_tokens.forward(input_ids)?) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_kv_cache() + } + } +} diff --git a/src/models/qwen3vl/config.rs b/src/models/qwen3vl/config.rs index 7de3280..91ee0d8 100644 --- a/src/models/qwen3vl/config.rs +++ b/src/models/qwen3vl/config.rs @@ -1,5 +1,7 @@ use candle_nn::Activation; +use crate::models::qwen3::config::Qwen3Config; + #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct Size { pub longest_edge: usize, @@ -46,6 +48,32 @@ pub struct Qwen3VLTextConfig { pub vocab_size: usize, } +pub fn qwen3vl_text_config2qwen3_config(cfg: &Qwen3VLTextConfig) -> Qwen3Config { + Qwen3Config { + attention_bias: cfg.attention_bias, + attention_dropout: cfg.attention_dropout as f64, + bos_token_id: cfg.bos_token_id as u32, + eos_token_id: cfg.eos_token_id as u32, + head_dim: cfg.head_dim, + hidden_act: cfg.hidden_act, + hidden_size: cfg.hidden_size, + initializer_range: cfg.initializer_range as f64, + intermediate_size: cfg.intermediate_size, + max_position_embeddings: cfg.max_position_embeddings, + max_window_layers: 0, + num_attention_heads: cfg.num_attention_heads, + num_hidden_layers: cfg.num_hidden_layers, + num_key_value_heads: cfg.num_key_value_heads, + rms_norm_eps: cfg.rms_norm_eps, + rope_theta: cfg.rope_theta, + tie_word_embeddings: true, + torch_dtype: cfg.dtype.clone(), + use_cache: cfg.use_cache, + use_sliding_window: false, + vocab_size: cfg.vocab_size, + } +} + #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct Qwen3VLVisionConfig { pub deepstack_visual_indexes: Vec, @@ -73,15 +101,3 @@ pub struct Qwen3VLConfig { pub vision_end_token_id: usize, pub vision_start_token_id: usize, } - -#[derive(Debug, Clone, PartialEq, serde::Deserialize)] -pub struct Qwen3VLGenerationConfig { - pub bos_token_id: usize, - pub pad_token_id: usize, - pub do_sample: bool, - pub eos_token_id: Vec, - pub top_p: f32, - pub top_k: usize, - pub temperature: f32, - pub repetition_penalty: f32, -} diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index fd7cedc..7b41f09 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -11,11 +11,8 @@ use crate::{ chat_template::ChatTemplate, models::{ GenerateModel, - qwen3vl::{ - config::{Qwen3VLConfig, Qwen3VLGenerationConfig}, - model::Qwen3VLModel, - processor::Qwen3VLProcessor, - }, + qwen3::config::Qwen3GenerationConfig, + qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor}, }, tokenizer::TokenizerModel, utils::{ @@ -32,7 +29,7 @@ pub struct Qwen3VLGenerateModel<'a> { device: Device, eos_token_id1: u32, eos_token_id2: u32, - generation_config: Qwen3VLGenerationConfig, + generation_config: Qwen3GenerationConfig, model_name: String, } @@ -50,7 +47,7 @@ impl<'a> Qwen3VLGenerateModel<'a> { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let qwen3_vl = Qwen3VLModel::new(cfg, vb)?; let generation_config_path = path.to_string() + "/generation_config.json"; - let generation_config: Qwen3VLGenerationConfig = + let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; Ok(Self { chat_template, diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index f0ec97a..6bea3d4 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -7,12 +7,14 @@ use candle_nn::{ use crate::{ models::{ - common::{GateUpDownMLP, TwoLinearMLP, eager_attention_forward, get_layer_norm}, - qwen3vl::config::{Qwen3VLConfig, Qwen3VLTextConfig, Qwen3VLVisionConfig}, + common::{TwoLinearMLP, eager_attention_forward, get_layer_norm}, + qwen3::model::Qwen3DecoderLayer, + qwen3vl::config::{ + Qwen3VLConfig, Qwen3VLTextConfig, Qwen3VLVisionConfig, qwen3vl_text_config2qwen3_config, + }, }, position_embed::rope::{ - Qwen2_5VisionRotaryEmbedding, Qwen3VLTextRotaryEmbedding, apply_rotary_pos_emb, - apply_rotary_pos_emb_vision, + Qwen2_5VisionRotaryEmbedding, Qwen3VLTextRotaryEmbedding, apply_rotary_pos_emb_vision, }, utils::tensor_utils::{ bitor_tensor, get_vision_next_indices, linspace, mask_index_add, masked_scatter_dim0, @@ -519,181 +521,9 @@ impl Qwen3VLVisionModel { } } -pub struct Qwen3VLTextAttention { - q_proj: Linear, - k_proj: Linear, - v_proj: Linear, - o_proj: Linear, - q_norm: RmsNorm, - k_norm: RmsNorm, - num_attention_heads: usize, - num_key_value_heads: usize, - num_kv_groups: usize, - head_dim: usize, - scaling: f64, - kv_cache: Option<(Tensor, Tensor)>, -} - -impl Qwen3VLTextAttention { - pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result { - let hidden_size = config.hidden_size; - let num_attention_heads = config.num_attention_heads; - let head_dim = config.head_dim; - let num_key_value_heads = config.num_key_value_heads; - let num_kv_groups = num_attention_heads / num_key_value_heads; - let scaling = 1f64 / f64::sqrt(head_dim as f64); - let (q_proj, k_proj, v_proj, o_proj) = if config.attention_bias { - let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?; - let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?; - let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?; - let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?; - (q_proj, k_proj, v_proj, o_proj) - } else { - let q_proj = - linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?; - let k_proj = - linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?; - let v_proj = - linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?; - let o_proj = - linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?; - (q_proj, k_proj, v_proj, o_proj) - }; - let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?; - let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?; - Ok(Self { - q_proj, - k_proj, - v_proj, - o_proj, - q_norm, - k_norm, - num_attention_heads, - num_key_value_heads, - num_kv_groups, - head_dim, - scaling, - kv_cache: None, - }) - } - - pub fn forward( - &mut self, - xs: &Tensor, - cos: &Tensor, - sin: &Tensor, - attention_mask: Option<&Tensor>, - ) -> Result { - let (b_sz, q_len, _) = xs.dims3()?; - let query_states = self.q_proj.forward(xs)?.reshape(( - b_sz, - q_len, - self.num_attention_heads, - self.head_dim, - ))?; - let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?; - let key_states = self.k_proj.forward(xs)?.reshape(( - b_sz, - q_len, - self.num_key_value_heads, - self.head_dim, - ))?; - let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?; - let value_states = self.v_proj.forward(xs)?; - let value_states = value_states - .reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))? - .transpose(1, 2)?; - let (query_states, key_states) = - apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?; - let (key_states, value_states) = match &self.kv_cache { - None => (key_states, value_states), - Some((prev_k, prev_v)) => { - let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; - let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; - (key_states, value_states) - } - }; - self.kv_cache = Some((key_states.clone(), value_states.clone())); - let attn_output = eager_attention_forward( - &query_states, - &key_states, - &value_states, - Some(self.num_kv_groups), - attention_mask, - self.scaling, - )?; - let attn_output = - attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?; - let attn_output = attn_output.apply(&self.o_proj)?; - Ok(attn_output) - } - - pub fn clear_kv_cache(&mut self) { - self.kv_cache = None - } -} - -pub struct Qwen3VLTextDecoderLayer { - self_attn: Qwen3VLTextAttention, - mlp: GateUpDownMLP, - input_layernorm: RmsNorm, - post_attention_layernorm: RmsNorm, -} - -impl Qwen3VLTextDecoderLayer { - pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result { - let self_attn = Qwen3VLTextAttention::new(config.clone(), vb.pp("self_attn"))?; - let mlp = GateUpDownMLP::new( - vb.pp("mlp"), - config.hidden_size, - config.intermediate_size, - config.hidden_act, - false, - )?; - let input_layernorm = rms_norm( - config.hidden_size, - config.rms_norm_eps, - vb.pp("input_layernorm"), - )?; - let post_attention_layernorm = rms_norm( - config.hidden_size, - config.rms_norm_eps, - vb.pp("post_attention_layernorm"), - )?; - Ok(Self { - self_attn, - mlp, - input_layernorm, - post_attention_layernorm, - }) - } - - pub fn forward( - &mut self, - xs: &Tensor, - cos: &Tensor, - sin: &Tensor, - attention_mask: Option<&Tensor>, - ) -> Result { - let residual = xs.clone(); - let xs = self.input_layernorm.forward(xs)?; - let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?; - let xs = residual.add(&xs)?; - let residual = xs.clone(); - let xs = self.post_attention_layernorm.forward(&xs)?; - let xs = self.mlp.forward(&xs)?; - let xs = residual.add(&xs)?; - Ok(xs) - } - - pub fn clear_kv_cache(&mut self) { - self.self_attn.clear_kv_cache(); - } -} - pub struct Qwen3VLTextModel { embed_tokens: Embedding, - layers: Vec, + layers: Vec, norm: RmsNorm, rotary_emb: Qwen3VLTextRotaryEmbedding, mrope_section: Vec, @@ -706,7 +536,8 @@ impl Qwen3VLTextModel { let mut layers = vec![]; let vb_l = vb.pp("layers"); for layer_idx in 0..config.num_hidden_layers { - let layer = Qwen3VLTextDecoderLayer::new(config.clone(), vb_l.pp(layer_idx))?; + let qwen3_cfg = qwen3vl_text_config2qwen3_config(&config); + let layer = Qwen3DecoderLayer::new(&qwen3_cfg, vb_l.pp(layer_idx))?; layers.push(layer) } let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?; diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 0071a81..302dae2 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -147,6 +147,7 @@ impl VoxCPMGenerate { } None => self.generate_simple(target_text)?, }; + self.voxcpm.clear_kv_cache(); Ok(audio) } @@ -197,6 +198,7 @@ impl VoxCPMGenerate { // retry_badcase, retry_badcase_ratio_threshold, )?; + self.voxcpm.clear_kv_cache(); Ok(audio) } @@ -223,19 +225,25 @@ impl GenerateModel for VoxCPMGenerate { } else { None }; - let audio = self.voxcpm.generate( - target_text, - prompt_text, - prompt_wav_path, - min_len, - max_len, - inference_timesteps, - cfg_value, - retry_badcase_ratio_threshold, - )?; + let audio = self + .voxcpm + .generate( + target_text, + prompt_text, + prompt_wav_path, + min_len, + max_len, + inference_timesteps, + cfg_value, + retry_badcase_ratio_threshold, + ) + .inspect_err(|_| { + self.voxcpm.clear_kv_cache(); + })?; let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate as u32)?; let base64_audio = BASE64_STANDARD.encode(wav_u8); let response = build_audio_completion_response(&base64_audio, &self.model_name); + self.voxcpm.clear_kv_cache(); Ok(response) } #[allow(unused_variables)] diff --git a/src/models/voxcpm/minicpm4.rs b/src/models/voxcpm/minicpm4.rs index 45529a4..756837e 100644 --- a/src/models/voxcpm/minicpm4.rs +++ b/src/models/voxcpm/minicpm4.rs @@ -120,6 +120,9 @@ impl MiniCPMDecoderLayer { None, false, None, + None, + None, + None, )?; let mlp = GateUpDownMLP::new( vb.pp("mlp"), diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs index eaacad3..d90ff7a 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -716,9 +716,15 @@ impl VoxCPMModel { .permute((0, 3, 1, 2))? .reshape((b, d, ()))? .contiguous()?; + // self.base_lm.clear_kv_cache(); + // self.residual_lm.clear_kv_cache(); + self.clear_kv_cache(); + Ok(feat_pred) + } + + pub fn clear_kv_cache(&mut self) { self.base_lm.clear_kv_cache(); self.residual_lm.clear_kv_cache(); - Ok(feat_pred) } pub fn build_prompt_cache( diff --git a/src/position_embed/mod.rs b/src/position_embed/mod.rs index f67f4f0..b0b3921 100644 --- a/src/position_embed/mod.rs +++ b/src/position_embed/mod.rs @@ -1 +1,2 @@ pub mod rope; +pub mod sinusoidal_pe; diff --git a/src/position_embed/sinusoidal_pe.rs b/src/position_embed/sinusoidal_pe.rs new file mode 100644 index 0000000..da0fcbb --- /dev/null +++ b/src/position_embed/sinusoidal_pe.rs @@ -0,0 +1,59 @@ +use anyhow::Result; +use candle_core::{D, DType, Device, Tensor}; + +use crate::position_embed::rope::compute_default_rope_parameters; + +pub struct SinusoidalPositionEncoderCat { + inv_freq: Option, // (1, dim / 2) +} + +impl SinusoidalPositionEncoderCat { + pub fn new(dim: Option, save_freq: bool, device: &Device) -> Result { + let inv_freq = if save_freq && let Some(dim) = dim { + let inv_freq = compute_default_rope_parameters(dim, 10000.0); + let inv_freq = Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)?; + Some(inv_freq) + } else { + None + }; + + Ok(Self { inv_freq }) + } + pub fn encode( + &self, + seqlen_offset: usize, + seq_len: usize, + head_dim: usize, + device: &Device, + dtype: DType, + ) -> Result { + let positions = Tensor::arange( + seqlen_offset as f32, + (seqlen_offset + seq_len) as f32, + device, + )? + .reshape((seq_len, 1))?; // (seq_len, 1) + let inv_freq = if self.inv_freq.is_none() { + let inv_freq = compute_default_rope_parameters(head_dim, 10000.0); + Tensor::from_slice(&inv_freq, (1, inv_freq.len()), device)? + } else { + self.inv_freq.as_ref().unwrap().clone() + }; + let freqs = positions.matmul(&inv_freq)?; // (seq_len, dim / 2) + let sin = freqs.sin()?; + let cos = freqs.cos()?; + let pos_embed = Tensor::cat(&[sin, cos], D::Minus1)? + .contiguous()? + .to_dtype(dtype)?; + + Ok(pos_embed) + } + pub fn forward(&self, xs: &Tensor, seqlen_offset: usize) -> Result { + let (_, seq_len, head_dim) = xs.dims3()?; + let pos_embed = self + .encode(seqlen_offset, seq_len, head_dim, xs.device(), xs.dtype())? + .unsqueeze(0)?; + let xs = xs.broadcast_add(&pos_embed)?; + Ok(xs) + } +} diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index c98d91e..0a569b3 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -1,6 +1,7 @@ use std::fs::File; use std::io::Write; use std::path::{Path, PathBuf}; +use std::thread; use std::{f64::consts::PI, io::Cursor}; use aha_openai_dive::v1::resources::chat::{ @@ -31,7 +32,7 @@ use symphonia::core::meta::MetadataOptions; use symphonia::core::probe::Hint; use crate::utils::get_default_save_dir; -use crate::utils::tensor_utils::linspace; +use crate::utils::tensor_utils::{linspace, pad_replicate_last_dim}; // 重采样方法枚举 #[derive(Debug, Clone, Copy)] @@ -279,19 +280,39 @@ pub fn load_audio_from_url(url: &str) -> Result { }) } +// pub fn _load_audio_bytes_from_url(url: &str) -> Result> { +// tokio::task::block_in_place(|| { +// let client = reqwest::blocking::Client::new(); +// let response = client.get(url).send()?; +// if !response.status().is_success() { +// return Err(anyhow::anyhow!( +// "Failed to download file: {}", +// response.status() +// )); +// } +// let bytes = response.bytes()?.to_vec(); +// Ok(bytes) +// }) +// } + pub fn load_audio_bytes_from_url(url: &str) -> Result> { - tokio::task::block_in_place(|| { - let client = reqwest::blocking::Client::new(); - let response = client.get(url).send()?; - if !response.status().is_success() { - return Err(anyhow::anyhow!( - "Failed to download file: {}", - response.status() - )); - } - let bytes = response.bytes()?.to_vec(); - Ok(bytes) + let url = url.to_string(); + thread::spawn(move || { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let response = reqwest::get(&url).await?; + if !response.status().is_success() { + return Err(anyhow::anyhow!( + "Failed to download file: {}", + response.status() + )); + } + let bytes = response.bytes().await?.to_vec(); + Ok(bytes) + }) }) + .join() + .unwrap() } pub fn get_audio_path(path_str: &str) -> Result { @@ -936,6 +957,31 @@ pub fn create_hann_window(window_size: usize, dtype: DType, device: &Device) -> Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?) } +pub fn crate_hamming_window( + window_size: usize, + periodic: bool, + alpha: f64, + beta: f64, + dtype: DType, + device: &Device, +) -> Result { + let denominator = if periodic { + window_size as f64 + } else { + (window_size - 1) as f64 + }; + + let window: Vec = (0..window_size) + .map(|i| { + let i_f64 = i as f64; + let val = alpha - beta * (2.0 * std::f64::consts::PI * i_f64 / denominator).cos(); + val as f32 + }) + .collect(); + + Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?) +} + /// 梅尔频率刻度类型 #[derive(Debug, Clone, Copy)] pub enum MelScale { @@ -1099,3 +1145,347 @@ pub fn stft_audio(n_fft: usize, frame_wave: &[f32]) -> Result> { let output: Vec = spectrum.iter().map(|complex| complex.norm_sqr()).collect(); Ok(output) } + +pub fn apply_stft(waveform: &Tensor) -> Result { + // waveform: (bs, n_frames, window_size) + let mut wave_fft = vec![]; + let (batch_size, _, window_size) = waveform.dims3()?; + for bs in 0..batch_size { + let wave_i = waveform.i(bs)?; + let wave_i_vec = wave_i.to_vec2::()?; + let wave_i_fft_vec: Result>> = wave_i_vec + .par_iter() + .map(|frame_wave| stft_audio(window_size, frame_wave)) + .collect(); + let wave_i_fft_vec = wave_i_fft_vec?; + + let wave_i_fft = Tensor::new(wave_i_fft_vec, waveform.device())?.unsqueeze(0)?; + wave_fft.push(wave_i_fft); + } + let magnitudes = Tensor::cat(&wave_fft, 0)?; + Ok(magnitudes) +} + +pub fn kaldi_fbank( + waveform: &Tensor, + mel_energies: &Tensor, + window_shift: usize, + window_size: usize, + padded_window_size: usize, + dither: f32, + // energy_floor: f32, + // window_type: &str, + // sample_frequency: usize, + // snip_edges: bool, +) -> Result { + let (strided_input, _) = get_window( + waveform, + padded_window_size, + window_size, + window_shift, + dither, + true, + true, + 0.97, + )?; + + let spectrum = apply_stft(&strided_input)?; + let mel_energies = spectrum.broadcast_matmul(mel_energies)?; + let epsilon = + Tensor::new(1.192_092_9e-7_f32, waveform.device())?.broadcast_as(mel_energies.shape())?; + let mel_energies = mel_energies.maximum(&epsilon)?.log()?; + + Ok(mel_energies) +} + +pub fn apply_lfr(inputs: &Tensor, lfr_m: usize, lfr_n: usize) -> Result { + let (t, feat_dim) = inputs.dims2()?; + let t_lfr = (t as f32 / lfr_n as f32).ceil() as usize; + let left_padding_size = (lfr_m - 1) / 2; + let left_padding = inputs.narrow(0, 0, 1)?.repeat((left_padding_size, 1))?; + let mut inputs = Tensor::cat(&[&left_padding, inputs], 0)?; + let t = t + left_padding_size; + let last_idx = (t - lfr_m) / lfr_n + 1; + let num_padding = lfr_m - (t - last_idx * lfr_n); + if num_padding > 0 { + let num_padding = + (2 * lfr_m - 2 * t + (t_lfr - 1 + last_idx) * lfr_n) / 2 * (t_lfr - last_idx); + let right_padding = inputs.narrow(0, t - 1, 1)?.repeat((num_padding, 1))?; + inputs = Tensor::cat(&[&inputs, &right_padding], 0)?; + } + let mut outputs = vec![]; + for i in 0..t_lfr { + let start = i * lfr_n; + let frame = inputs + .narrow(0, start, lfr_m)? + .reshape((1, lfr_m * feat_dim))?; + outputs.push(frame); + } + let lfr_outputs = Tensor::cat(&outputs, 0)?; + Ok(lfr_outputs) +} + +pub fn get_waveform_and_window_properties( + sample_frequency: usize, + frame_shift: f32, + frame_length: f32, + round_to_power_of_two: bool, +) -> Result<(usize, usize, usize)> { + let window_shift = (sample_frequency as f32 * frame_shift * 0.001) as usize; + let window_size = (sample_frequency as f32 * frame_length * 0.001) as usize; + let padded_window_size = if round_to_power_of_two { + (window_size - 1).next_power_of_two() + } else { + window_size + }; + Ok((window_shift, window_size, padded_window_size)) +} + +pub fn get_window( + waveform: &Tensor, + padded_window_size: usize, + window_size: usize, + window_shift: usize, + dither: f32, + remove_dc_offset: bool, + raw_energy: bool, + preemphasis_coefficient: f32, +) -> Result<(Tensor, Tensor)> { + let mut strided_input = extract_frames(waveform, window_size, window_shift)?; + // (ba, m, window_size) + if dither != 0.0 { + let rand_gauss = strided_input + .randn_like(0.0, 1.0)? + .affine(dither as f64, 0.0)?; + strided_input = strided_input.add(&rand_gauss)?; + } + if remove_dc_offset { + let row_means = strided_input.mean_keepdim(D::Minus1)?; + strided_input = strided_input.broadcast_sub(&row_means)?; + } + let signal_log_energy = if raw_energy { + let energy = strided_input.powf(2.0)?.sum(1)?.log()?; + Some(energy) + } else { + None + }; + + if preemphasis_coefficient != 0.0 { + let offset_strided_input = pad_replicate_last_dim(&strided_input, (1, 0))? + .affine(preemphasis_coefficient as f64, 0.0)?; + strided_input = + strided_input.sub(&offset_strided_input.narrow(D::Minus1, 0, window_size)?)?; + } + + let windows = crate_hamming_window( + window_size, + false, + 0.54, + 0.46, + waveform.dtype(), + waveform.device(), + )? + .unsqueeze(0)? + .unsqueeze(0)?; + + strided_input = strided_input.broadcast_mul(&windows)?; + + if padded_window_size != window_size { + let padding_right = padded_window_size - window_size; + strided_input = strided_input.pad_with_zeros(D::Minus1, 0, padding_right)?; + } + + let signal_log_energy = signal_log_energy.unwrap_or(strided_input.powf(2.0)?.sum(1)?.log()?); + Ok((strided_input, signal_log_energy)) +} + +/// 提取音频帧 +pub fn extract_frames( + waveform: &Tensor, + window_size: usize, + window_shift: usize, +) -> Result { + // waveform ->(1, audio_len) + let waveform_len = waveform.dim(1)?; + let n_frames = 1 + (waveform_len - window_size) / window_shift; + let mut frames = Vec::with_capacity(n_frames); + + for i in 0..n_frames { + let start = i * window_shift; + let frame = waveform.narrow(D::Minus1, start, window_size)?; + frames.push(frame); + } + + let result = Tensor::cat(&frames, D::Minus1)?; + let bs = result.dim(0)?; + let reshaped = result.reshape((bs, n_frames, window_size))?; + Ok(reshaped) +} + +pub fn inverse_mel_scale(mel_freq: &Tensor) -> Result { + Ok(mel_freq + .affine(1.0 / 1127.0, 0.0)? + .exp()? + .affine(1.0, -1.0)? + .affine(700.0, 0.0)?) +} + +pub fn mel_scale(freq: &Tensor) -> Result { + Ok(freq.affine(1.0 / 700.0, 1.0)?.log()?.affine(1127.0, 0.0)?) +} + +pub fn kaldi_get_mel_banks( + num_bins: usize, + window_length_padded: usize, + sample_freq: f32, + low_freq: f32, + high_freq: f32, + // vtln_low: f32, + // vtln_high: f32, + // vtln_warp_factor: f32, + device: &Device, +) -> Result<(Tensor, Tensor)> { + assert!(num_bins > 3, "Must have at least 3 mel bins"); + assert!( + window_length_padded.is_multiple_of(2), + "window_length_padded must be even" + ); + + let num_fft_bins = window_length_padded as f32 / 2.0; + let nyquist = 0.5 * sample_freq; + + let mut high_freq = high_freq; + if high_freq <= 0.0 { + high_freq += nyquist; + } + + assert!( + (0.0 <= low_freq && low_freq < nyquist) + && (0.0 < high_freq && high_freq <= nyquist) + && (low_freq < high_freq), + "Bad values in options: low-freq {} and high-freq {} vs. nyquist {}", + low_freq, + high_freq, + nyquist + ); + + // FFT bin 宽度 + let fft_bin_width = sample_freq / (window_length_padded as f32); + let mel_low_freq = hertz_to_mel(low_freq, MelScale::Kaldi); + let mel_high_freq = hertz_to_mel(high_freq, MelScale::Kaldi); + + // 分频点之间的间隔 + let mel_freq_delta = (mel_high_freq - mel_low_freq) / ((num_bins + 1) as f32); + + // let mut vtln_high = vtln_high; + // if vtln_high < 0.0 { + // vtln_high += nyquist; + // } + + // if vtln_warp_factor != 1.0 { + // assert!( + // low_freq < vtln_low + // && vtln_low < high_freq + // && 0.0 < vtln_high + // && vtln_high < high_freq + // && vtln_low < vtln_high, + // "Bad values in options: vtln-low {} and vtln-high {}, versus low-freq {} and high-freq {}", + // vtln_low, + // vtln_high, + // low_freq, + // high_freq + // ); + // } + + // 创建 bin 索引张量 + let bins = Tensor::arange(0u32, num_bins as u32, device)? + .to_dtype(candle_core::DType::F32)? + .unsqueeze(1)?; // size(num_bins, 1) + + // 计算梅尔刻度下的边界频率 + let left_mel = bins.affine(mel_freq_delta as f64, mel_low_freq as f64)?; + let center_mel = bins + .affine(1.0, 1.0)? + .affine(mel_freq_delta as f64, mel_low_freq as f64)?; + let right_mel = bins + .affine(1.0, 2.0)? + .affine(mel_freq_delta as f64, mel_low_freq as f64)?; + + // 如果使用 VTLN,则对频率进行扭曲 + // let (left_mel, center_mel, right_mel) = if vtln_warp_factor != 1.0 { + // ( + // vtln_warp_mel_freq( + // vtln_low, + // vtln_high, + // low_freq, + // high_freq, + // vtln_warp_factor, + // &left_mel, + // )?, + // vtln_warp_mel_freq( + // vtln_low, + // vtln_high, + // low_freq, + // high_freq, + // vtln_warp_factor, + // ¢er_mel, + // )?, + // vtln_warp_mel_freq( + // vtln_low, + // vtln_high, + // low_freq, + // high_freq, + // vtln_warp_factor, + // &right_mel, + // )?, + // ) + // } else { + // (left_mel, center_mel, right_mel) + // }; + + // 转换中心频率回赫兹单位 + let center_freqs = inverse_mel_scale(¢er_mel)?; + + // 创建 FFT bin 频率 + let fft_bins = Tensor::arange(0u32, num_fft_bins as u32, device)? + .to_dtype(candle_core::DType::F32)? + .affine(fft_bin_width as f64, 0.0)?; + let mel = mel_scale(&fft_bins)?.unsqueeze(0)?; // size(1, num_fft_bins) + + // 计算斜率 + let up_slope = mel + .broadcast_sub(&left_mel)? + .broadcast_div(¢er_mel.broadcast_sub(&left_mel)?)?; + let down_slope = right_mel + .broadcast_sub(&mel)? + .broadcast_div(&right_mel.broadcast_sub(¢er_mel)?)?; + + // left_mel < center_mel < right_mel 所以我们可以取两个斜率的最小值并限制负值 + let min_slopes = up_slope.minimum(&down_slope)?; + let zeros = Tensor::zeros(min_slopes.dims(), candle_core::DType::F32, device)?; + let bins_tensor = min_slopes.maximum(&zeros)?; + // let bins_tensor = if vtln_warp_factor == 1.0 { + // // left_mel < center_mel < right_mel 所以我们可以取两个斜率的最小值并限制负值 + // let min_slopes = up_slope.minimum(&down_slope)?; + // let zeros = Tensor::zeros(min_slopes.dims(), candle_core::DType::F32, device)?; + // min_slopes.maximum(&zeros)? + // } else { + // // 扭曲可能会改变 left_mel, center_mel, right_mel 的顺序 + // let zeros = Tensor::zeros(up_slope.dims(), candle_core::DType::F32, device)?; + // let mut bins_tensor = zeros.clone(); + + // // 创建索引掩码 + // let up_idx = mel + // .gt_tensor(&left_mel)? + // .and(&mel.le_tensor(¢er_mel)?)?; // left_mel < mel <= center_mel + // let down_idx = mel + // .gt_tensor(¢er_mel)? + // .and(&mel.lt_tensor(&right_mel)?)?; // center_mel < mel < right_mel + + // bins_tensor = bins_tensor.where_cond(&up_idx, &up_slope)?; + // bins_tensor = bins_tensor.where_cond(&down_idx, &down_slope)?; + // bins_tensor + // }; + + Ok((bins_tensor, center_freqs)) +} diff --git a/src/utils/img_utils.rs b/src/utils/img_utils.rs index 78838e4..d78cbf1 100644 --- a/src/utils/img_utils.rs +++ b/src/utils/img_utils.rs @@ -1,4 +1,5 @@ use std::io::Cursor; +use std::thread; use std::{collections::HashSet, path::PathBuf}; use aha_openai_dive::v1::resources::chat::{ @@ -13,22 +14,45 @@ use rayon::iter::{IntoParallelRefIterator, ParallelIterator}; use crate::utils::{ceil_by_factor, floor_by_factor, round_by_factor}; pub fn load_image_from_url(url: &str) -> Result { - tokio::task::block_in_place(|| { - let response = reqwest::blocking::get(url) - .map_err(|e| anyhow!(format!("Failed to fetch image from url: {}", e)))?; - let bytes = response - .bytes() - .map_err(|e| anyhow!(format!("Failed to get image bytes: {}", e)))?; + let url = url.to_string(); + thread::spawn(move || { + let rt = tokio::runtime::Runtime::new().unwrap(); + rt.block_on(async { + let response = reqwest::blocking::get(url) + .map_err(|e| anyhow!(format!("Failed to fetch image from url: {}", e)))?; + let bytes = response + .bytes() + .map_err(|e| anyhow!(format!("Failed to get image bytes: {}", e)))?; - let cursor = Cursor::new(bytes); - let img = ImageReader::new(cursor) - .with_guessed_format() - .map_err(|e| anyhow!(format!("Failed to read image format: {}", e)))? - .decode() - .map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?; - Ok(img) + let cursor = Cursor::new(bytes); + let img = ImageReader::new(cursor) + .with_guessed_format() + .map_err(|e| anyhow!(format!("Failed to read image format: {}", e)))? + .decode() + .map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?; + Ok(img) + }) }) + .join() + .unwrap() } +// pub fn _load_image_from_url(url: &str) -> Result { +// tokio::task::block_in_place(|| { +// let response = reqwest::blocking::get(url) +// .map_err(|e| anyhow!(format!("Failed to fetch image from url: {}", e)))?; +// let bytes = response +// .bytes() +// .map_err(|e| anyhow!(format!("Failed to get image bytes: {}", e)))?; + +// let cursor = Cursor::new(bytes); +// let img = ImageReader::new(cursor) +// .with_guessed_format() +// .map_err(|e| anyhow!(format!("Failed to read image format: {}", e)))? +// .decode() +// .map_err(|e| anyhow!(format!("Failed to decode image: {}", e)))?; +// Ok(img) +// }) +// } pub fn load_image_from_base64(base64_data: &str) -> Result { let image_data = general_purpose::STANDARD diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index 2ee4a2b..3ff8b42 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -2,6 +2,20 @@ use anyhow::{Result, anyhow}; use candle_core::{D, DType, Device, IndexOp, Tensor, shape::Dim}; use candle_nn::ops::sigmoid; +pub fn mask_filled(on_true: &Tensor, mask: &Tensor, on_false: f32) -> Result { + let (mask_seq_len, _) = mask.dims2()?; + let (_, _, seq_len, _) = on_true.dims4()?; + assert!( + mask_seq_len >= seq_len, + "mask seq_len less than input data seq_len" + ); + let mask = mask.i((..seq_len, ..seq_len))?; + let mask = mask.broadcast_as(on_true.shape())?; + let on_false = Tensor::new(on_false, on_true.device())?.broadcast_as(on_true.shape())?; + let filled = mask.where_cond(on_true, &on_false)?; + Ok(filled) +} + pub fn prepare_causal_attention_mask( b_size: usize, tgt_len: usize, @@ -870,3 +884,28 @@ pub fn pad_reflect_last_dim(t: &Tensor, pad: (usize, usize)) -> Result { } Ok(pad_tensor) } + +pub fn pad_replicate_last_dim(t: &Tensor, pad: (usize, usize)) -> Result { + let (pad_l, pad_r) = pad; + let last_dim = t.dim(D::Minus1)?; + + let mut pad_tensor = t.clone(); + if pad_l > 0 { + let left = pad_tensor.narrow(D::Minus1, 0, 1)?.contiguous()?; + let rank = left.rank(); + let mut shape = vec![1usize; rank - 1]; + shape.push(pad_l); + let left_pad = left.repeat(shape)?; + pad_tensor = Tensor::cat(&[&left_pad, &pad_tensor], D::Minus1)?; + } + if pad_r > 0 { + let start_i = last_dim - 1; + let right = pad_tensor.narrow(D::Minus1, start_i, 1)?.contiguous()?; + let rank = right.rank(); + let mut shape = vec![1usize; rank - 1]; + shape.push(pad_r); + let right_pad = right.repeat(shape)?; + pad_tensor = Tensor::cat(&[&pad_tensor, &right_pad], D::Minus1)?; + } + Ok(pad_tensor) +} diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs new file mode 100644 index 0000000..a24001e --- /dev/null +++ b/tests/test_fun_asr_nano.rs @@ -0,0 +1,97 @@ +use std::{pin::pin, time::Instant}; + +use aha::models::{GenerateModel, fun_asr_nano::generate::FunAsrNanoGenerateModel}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use rocket::futures::StreamExt; +#[test] +fn fun_asr_nano_generate() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda fun_asr_nano_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!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", save_dir); + let message = r#" + { + "model": "fun-asr-nano", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" + } + }, + { + "type": "text", + "text": "语音转写:" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + let i_start = Instant::now(); + let res = fun_asr_model.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("generate: \n {:?}", res); + if res.usage.is_some() { + let num_token = res.usage.as_ref().unwrap().total_tokens; + let duration_secs = i_duration.as_secs_f64(); + let tps = num_token as f64 / duration_secs; + println!("Tokens per second (TPS): {:.2}", tps); + } + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} + +#[tokio::test] +async fn fun_asr_nano_stream() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda fun_asr_nano_stream -r -- --nocapture + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", save_dir); + let message = r#" + { + "model": "fun-asr-nano", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" + } + }, + { + "type": "text", + "text": "语音转写:" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + let i_start = Instant::now(); + let mut stream = pin!(fun_asr_model.generate_stream(mes)?); + while let Some(item) = stream.next().await { + println!("generate: \n {:?}", item); + } + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs index db76795..78238d1 100644 --- a/tests/test_glm_asr_nano.rs +++ b/tests/test_glm_asr_nano.rs @@ -70,7 +70,7 @@ async fn glm_asr_nano_stream() -> Result<()> { "type": "audio", "audio_url": { - "url": "file://./assets/audio/zh.mp3" + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" } }, { diff --git a/tests/test_paddleocr_vl.rs b/tests/test_paddleocr_vl.rs index 682afb3..dbab38e 100644 --- a/tests/test_paddleocr_vl.rs +++ b/tests/test_paddleocr_vl.rs @@ -19,7 +19,7 @@ fn paddleocr_vl_generate() -> Result<()> { "type": "image", "image_url": { - "url": "file://./assets/img/ocr_test1.png" + "url": "https://www.qqxiuzi.cn/zh/shouxie-shufa/welcome.png" } }, { @@ -69,7 +69,7 @@ async fn paddleocr_vl_stream() -> Result<()> { "type": "image", "image_url": { - "url": "file://./assets/img/ocr_test1.png" + "url": "https://www.qqxiuzi.cn/zh/shouxie-shufa/welcome.png" } }, { diff --git a/tests/test_qwen3.rs b/tests/test_qwen3.rs new file mode 100644 index 0000000..2ba8b41 --- /dev/null +++ b/tests/test_qwen3.rs @@ -0,0 +1,83 @@ +use std::{pin::pin, time::Instant}; + +use aha::models::{GenerateModel, qwen3::generate::Qwen3GenerateModel}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use rocket::futures::StreamExt; + +#[test] +fn qwen3_0_6b_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3_0_6b_generate -r -- --nocapture + // test with cuda+flash-attn: RUST_BACKTRACE=1 cargo test -F cuda,flash-attn qwen3_0_6b_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!("{}/Qwen/Qwen3-0.6B/", save_dir); + let message = r#" + { + "model": "qwen3", + "messages": [ + { + "role": "user", + "content": "你吃饭了没" + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut model = Qwen3GenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let result = model.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("generate: \n {:?}", result); + if result.usage.is_some() { + let num_token = result.usage.as_ref().unwrap().total_tokens; + let duration_secs = i_duration.as_secs_f64(); + let tps = num_token as f64 / duration_secs; + println!("Tokens per second (TPS): {:.2}", tps); + } + println!("Time elapsed in generate is: {:?}", i_duration); + + Ok(()) +} + +#[tokio::test] +async fn qwen3_0_6b_stream() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3_0_6b_stream -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/Qwen/Qwen3-0.6B/", save_dir); + + let message = r#" + { + "model": "qwen3", + "messages": [ + { + "role": "user", + "content": "你是谁" + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut model = Qwen3GenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let mut stream = pin!(model.generate_stream(mes)?); + while let Some(item) = stream.next().await { + println!("generate: \n {:?}", item); + } + + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + Ok(()) +} diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index f52d662..9c81761 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -29,7 +29,7 @@ fn qwen3vl_generate() -> Result<()> { }, { "type": "text", - "text": "视频中发生了什么?, 现在几点了" + "text": "视频中发生了什么?" } ] } diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index c45464b..bdbd458 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -27,7 +27,7 @@ fn voxcpm_use_message_generate() -> Result<()> { "type": "audio", "audio_url": { - "url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.wav" + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" } }, { @@ -37,7 +37,7 @@ fn voxcpm_use_message_generate() -> Result<()> { ] } ], - "metadata": {"prompt_text": "华为致力于把数字世界带给每个人,每个家庭,每个组织,构建万物互联的智能世界。"} + "metadata": {"prompt_text": "天雷滚滚我好怕怕,劈得我浑身掉渣渣。突破天劫我笑哈哈,逆天改命我吹喇叭,滴答滴答滴滴答"} } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index 4c7f452..397a2d1 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -27,17 +27,17 @@ fn voxcpm1_5_use_message_generate() -> Result<()> { "type": "audio", "audio_url": { - "url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.wav" + "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" } }, { "type": "text", - "text": "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech." + "text": "老大爷我来啦,红红火火恍恍惚惚" } ] } ], - "metadata": {"prompt_text": "华为致力于把数字世界带给每个人,每个家庭,每个组织,构建万物互联的智能世界。"} + "metadata": {"prompt_text": "天雷滚滚我好怕怕,劈得我浑身掉渣渣。突破天劫我笑哈哈,逆天改命我吹喇叭,滴答滴答滴滴答"} } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; @@ -72,7 +72,7 @@ fn voxcpm1_5_generate() -> Result<()> { let i_start = Instant::now(); // let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?; let generate = voxcpm_generate.inference( - "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(), + "老大爷我来啦,红红火火恍恍惚惚".to_string(), Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()), Some("file://./assets/audio/voice_01.wav".to_string()), // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 32158fc..4fabe06 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -152,3 +152,47 @@ fn glm_asr_nano_weight() -> Result<()> { println!("model_list: {:?}", model_list); Ok(()) } + +#[test] +fn fun_asr_nano_weight() -> Result<()> { + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", 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, Some("state_dict"))?; + // dtype = dict[0].1.dtype(); + for (k, v) in dict { + if k.contains("model") { + println!("key: {}, tensor shape: {:?}", k, v); + } + dict_to_hashmap.insert(k, v); + } + } + Ok(()) +} + +#[test] +fn qwen3_weight() -> Result<()> { + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/Qwen/Qwen3-0.6B/", save_dir); + let model_list = find_type_files(&model_path, "safetensors")?; + + let device = Device::Cpu; + for m in &model_list { + let weights = safetensors::load(m, &device)?; + for (key, tensor) in weights.iter() { + // if key.contains(".embed_tokens") { + // println!("=== {} === {:?}", key, tensor.shape()); + // } + println!("=== {} === {:?}", key, tensor.shape()); + } + } + println!("model_list: {:?}", model_list); + Ok(()) +}