diff --git a/Cargo.lock b/Cargo.lock index fdb881a..42320ea 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -19,7 +19,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" [[package]] name = "aha" -version = "0.1.6" +version = "0.1.7" dependencies = [ "aha_openai_dive", "anyhow", @@ -38,10 +38,12 @@ dependencies = [ "modelscope", "num", "rayon", + "realfft", "reqwest", "rocket", "serde", "serde_json", + "symphonia", "tokenizers", "tokio", "url", @@ -896,7 +898,7 @@ dependencies = [ "libc", "option-ext", "redox_users", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -1002,7 +1004,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -1040,6 +1042,12 @@ dependencies = [ "zune-inflate", ] +[[package]] +name = "extended" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af9673d8203fcb076b19dfd17e38b3d4ae9f44959416ea532ce72415a6020365" + [[package]] name = "fancy-regex" version = "0.13.0" @@ -1802,7 +1810,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.5.10", + "socket2 0.6.1", "system-configuration", "tokio", "tower-service", @@ -2066,7 +2074,7 @@ checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" dependencies = [ "hermit-abi", "libc", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2457,7 +2465,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -2812,6 +2820,15 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "primal-check" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc0d895b311e3af9902528fbb8f928688abbd95872819320517cc24ca6b2bd08" +dependencies = [ + "num-integer", +] + [[package]] name = "proc-macro-crate" version = "3.4.0" @@ -3095,6 +3112,15 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "realfft" +version = "3.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f821338fddb99d089116342c46e9f1fbf3828dba077674613e734e01d6ea8677" +dependencies = [ + "rustfft", +] + [[package]] name = "reborrow" version = "0.5.5" @@ -3345,6 +3371,20 @@ version = "2.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "357703d41365b4b27c590e3ed91eabb1b663f07c4c084095e60cbed4362dff0d" +[[package]] +name = "rustfft" +version = "6.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "21db5f9893e91f41798c88680037dba611ca6674703c1a18601b01a72c8adb89" +dependencies = [ + "num-complex", + "num-integer", + "num-traits", + "primal-check", + "strength_reduce", + "transpose", +] + [[package]] name = "rustix" version = "1.1.2" @@ -3355,7 +3395,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -3677,6 +3717,12 @@ version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a2eb9349b6444b326872e140eb1cf5e7c522154d69e7a0ffb0fb81c06b37543f" +[[package]] +name = "strength_reduce" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe895eb47f22e2ddd4dabc02bce419d2e643c8e3b585c78158b349195bc24d82" + [[package]] name = "strsim" version = "0.11.1" @@ -3689,6 +3735,152 @@ version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +[[package]] +name = "symphonia" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5773a4c030a19d9bfaa090f49746ff35c75dfddfa700df7a5939d5e076a57039" +dependencies = [ + "lazy_static", + "symphonia-bundle-flac", + "symphonia-bundle-mp3", + "symphonia-codec-adpcm", + "symphonia-codec-pcm", + "symphonia-codec-vorbis", + "symphonia-core", + "symphonia-format-mkv", + "symphonia-format-ogg", + "symphonia-format-riff", + "symphonia-metadata", +] + +[[package]] +name = "symphonia-bundle-flac" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c91565e180aea25d9b80a910c546802526ffd0072d0b8974e3ebe59b686c9976" +dependencies = [ + "log", + "symphonia-core", + "symphonia-metadata", + "symphonia-utils-xiph", +] + +[[package]] +name = "symphonia-bundle-mp3" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4872dd6bb56bf5eac799e3e957aa1981086c3e613b27e0ac23b176054f7c57ed" +dependencies = [ + "lazy_static", + "log", + "symphonia-core", + "symphonia-metadata", +] + +[[package]] +name = "symphonia-codec-adpcm" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2dddc50e2bbea4cfe027441eece77c46b9f319748605ab8f3443350129ddd07f" +dependencies = [ + "log", + "symphonia-core", +] + +[[package]] +name = "symphonia-codec-pcm" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e89d716c01541ad3ebe7c91ce4c8d38a7cf266a3f7b2f090b108fb0cb031d95" +dependencies = [ + "log", + "symphonia-core", +] + +[[package]] +name = "symphonia-codec-vorbis" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f025837c309cd69ffef572750b4a2257b59552c5399a5e49707cc5b1b85d1c73" +dependencies = [ + "log", + "symphonia-core", + "symphonia-utils-xiph", +] + +[[package]] +name = "symphonia-core" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea00cc4f79b7f6bb7ff87eddc065a1066f3a43fe1875979056672c9ef948c2af" +dependencies = [ + "arrayvec", + "bitflags 1.3.2", + "bytemuck", + "lazy_static", + "log", +] + +[[package]] +name = "symphonia-format-mkv" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "122d786d2c43a49beb6f397551b4a050d8229eaa54c7ddf9ee4b98899b8742d0" +dependencies = [ + "lazy_static", + "log", + "symphonia-core", + "symphonia-metadata", + "symphonia-utils-xiph", +] + +[[package]] +name = "symphonia-format-ogg" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b4955c67c1ed3aa8ae8428d04ca8397fbef6a19b2b051e73b5da8b1435639cb" +dependencies = [ + "log", + "symphonia-core", + "symphonia-metadata", + "symphonia-utils-xiph", +] + +[[package]] +name = "symphonia-format-riff" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2d7c3df0e7d94efb68401d81906eae73c02b40d5ec1a141962c592d0f11a96f" +dependencies = [ + "extended", + "log", + "symphonia-core", + "symphonia-metadata", +] + +[[package]] +name = "symphonia-metadata" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "36306ff42b9ffe6e5afc99d49e121e0bd62fe79b9db7b9681d48e29fa19e6b16" +dependencies = [ + "encoding_rs", + "lazy_static", + "log", + "symphonia-core", +] + +[[package]] +name = "symphonia-utils-xiph" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee27c85ab799a338446b68eec77abf42e1a6f1bb490656e121c6e27bfbab9f16" +dependencies = [ + "symphonia-core", + "symphonia-metadata", +] + [[package]] name = "syn" version = "2.0.108" @@ -3798,7 +3990,7 @@ dependencies = [ "getrandom 0.3.4", "once_cell", "rustix", - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] @@ -4187,6 +4379,16 @@ dependencies = [ "tracing-log", ] +[[package]] +name = "transpose" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ad61aed86bc3faea4300c7aee358b4c6d0c8d6ccc36524c96e4c92ccf26e77e" +dependencies = [ + "num-integer", + "strength_reduce", +] + [[package]] name = "try-lock" version = "0.2.5" @@ -4524,7 +4726,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.61.2", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index fb504d7..65f9fcb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,10 +1,10 @@ [package] name = "aha" -version = "0.1.6" +version = "0.1.7" 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" +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" [dependencies] candle-core = { version = "0.9.1"} @@ -32,6 +32,10 @@ modelscope = "0.1.0" dirs = "6.0.0" url = "2.5.7" rayon = "1.10" +# rubato = "1.0.0" +# audioadapter-buffers = "2.0.0" +realfft = "3.5.0" +symphonia = { version = "0.5.5", features = ["mp3", "wav"] } [features] flash-attn=["candle-flash-attn"] @@ -41,3 +45,4 @@ ffmpeg=["ffmpeg-next"] [lints.clippy] needless_range_loop = "allow" single_range_in_vec_init = "allow" +manual_div_ceil = "allow" diff --git a/README.md b/README.md index 8218c0e..68dfd13 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,8 @@ - 模型:[VoxCPM1.5](https://huggingface.co/openbmb/VoxCPM1.5) 开源协议:[Apache license 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) * [RMBG2.0](https://huggingface.co/collections/briaai/rmbg) - RMBGv2.0由BRIA AI开发,供非商业用途使用。 - 模型:[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) ## 计划支持 我们持续扩展支持的模型列表,欢迎贡献! @@ -114,6 +116,7 @@ cargo run -F cuda -r -- [参数] * RMBG2.0: AI-ModelScope/RMBG-2.0 模型 * voxcpm: OpenBMB/VoxCPM-0.5B 模型 * voxcpm1.5: OpenBMB/VoxCPM1.5 模型 + * glm-asr-nano-2512: ZhipuAI/GLM-ASR-Nano-2512 模型 * 示例:--model deepseek-ocr 或 -m qwen3vl-2b 3. 权重路径 @@ -150,7 +153,7 @@ cargo run -F cuda -r -- [参数] 1. 对话接口 - **端点**: `POST /chat/completions` - **功能**: 多模态对话和文本生成 -- **支持模型**: Qwen2.5VL,Qwen3VL,DeepSeekOCR 等 +- **支持模型**: Qwen2.5VL,Qwen3VL,DeepSeekOCR, GLM-ASR-Nano-2512 等 - **请求格式**: OpenAI Chat Completion 格式 - **响应格式**: OpenAI Chat Completion 格式 - **流式支持**: 支持 @@ -286,6 +289,9 @@ cargo test -F cuda voxcpm_generate -r -- --nocapture 2. 提交新的 Issue,包含详细描述和复现步骤 ## 更新日志 +### v0.1.7 +* 支持GLM-ASR-Nano-2512 模型 + ### v0.1.6 * 支持RMGB2.0 模型 diff --git a/assets/audio/zh.mp3 b/assets/audio/zh.mp3 new file mode 100644 index 0000000..1ae2c89 Binary files /dev/null and b/assets/audio/zh.mp3 differ diff --git a/src/main.rs b/src/main.rs index f08fd3b..2126c5d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -84,6 +84,7 @@ async fn main() -> anyhow::Result<()> { WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0", WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B", WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5", + WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-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 9782441..f72672e 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -1,12 +1,16 @@ use anyhow::Result; use candle_core::{D, Tensor}; use candle_nn::{ - Activation, BatchNorm, BatchNormConfig, Conv2d, Conv2dConfig, LayerNorm, LayerNormConfig, - Linear, Module, RmsNorm, VarBuilder, batch_norm, conv2d, conv2d_no_bias, layer_norm, linear, - linear_no_bias, rms_norm, + Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, Embedding, + LayerNorm, LayerNormConfig, Linear, Module, RmsNorm, VarBuilder, batch_norm, conv1d, + conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear, linear_no_bias, + rms_norm, }; -use crate::{position_embed::rope::apply_rotary_pos_emb, utils::tensor_utils::repeat_kv}; +use crate::{ + position_embed::rope::{RoPE, apply_rotary_pos_emb}, + utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv}, +}; #[derive(Debug, Clone)] pub struct GateUpDownMLP { @@ -63,8 +67,11 @@ pub struct TwoLinearMLP { impl TwoLinearMLP { pub fn new( vb: VarBuilder, - embedding_dim: usize, - mlp_dim: usize, + // embedding_dim: usize, + // mlp_dim: usize, + in_dim: usize, + middle_dim: usize, + out_dim: usize, act: Activation, bias: bool, linear1_pp_name: &str, @@ -72,13 +79,13 @@ impl TwoLinearMLP { ) -> Result { let (linear1, linear2) = if bias { ( - linear(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?, - linear(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?, + linear(in_dim, middle_dim, vb.pp(linear1_pp_name))?, + linear(middle_dim, out_dim, vb.pp(linear2_pp_name))?, ) } else { ( - linear_no_bias(embedding_dim, mlp_dim, vb.pp(linear1_pp_name))?, - linear_no_bias(mlp_dim, embedding_dim, vb.pp(linear2_pp_name))?, + linear_no_bias(in_dim, middle_dim, vb.pp(linear1_pp_name))?, + linear_no_bias(middle_dim, out_dim, vb.pp(linear2_pp_name))?, ) }; Ok(Self { @@ -304,6 +311,7 @@ impl NaiveAttnTwoLinearMLPBlock { vb.pp(mlp_pp_name), hidden_size, intermediate_size, + hidden_size, hidden_act, mlp_bias, linear1_pp_name, @@ -503,6 +511,32 @@ pub fn get_conv2d( Ok(conv2d) } +pub fn get_conv1d( + vb: VarBuilder, + in_c: usize, + out_c: usize, + kernel_size: usize, + padding: usize, + stride: usize, + dilation: usize, + groups: usize, + bias: bool, +) -> Result { + let cfg = Conv1dConfig { + padding, + stride, + dilation, + groups, + cudnn_fwd_algo: None, + }; + let conv1d = if bias { + conv1d(in_c, out_c, kernel_size, cfg, vb)? + } else { + conv1d_no_bias(in_c, out_c, kernel_size, cfg, vb)? + }; + Ok(conv1d) +} + pub fn get_layer_norm(vb: VarBuilder, eps: f64, dim: usize) -> Result { let ln_config = LayerNormConfig { eps, @@ -620,3 +654,160 @@ pub fn deform_conv2d_kernel( } Ok(out) } + +pub struct LlamaModel { + pub embed_tokens: Embedding, + layers: Vec, + norm: RmsNorm, + rotary_emb: RoPE, +} + +impl LlamaModel { + pub fn new( + vb: VarBuilder, + vocab_size: usize, + hidden_size: usize, + num_hidden_layers: usize, + num_attention_heads: usize, + num_key_value_heads: Option, + head_dim: Option, + attn_bias: bool, + attn_pp_name: &str, + o_proj_pp_name: Option<&str>, + intermediate_size: usize, + hidden_act: Activation, + mlp_bias: bool, + mlp_pp_name: &str, + norm_eps: f64, + input_norm_pp_name: &str, + post_norm_pp_name: &str, + rope_theta_base: f32, + ) -> Result { + let embed_tokens = embedding(vocab_size, hidden_size, vb.pp("embed_tokens"))?; + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + for i in 0..num_hidden_layers { + let layers_i = NaiveAttnGateUpDownMLPBlock::new( + vb_layers.pp(i), + hidden_size, + num_attention_heads, + num_key_value_heads, + head_dim, + attn_bias, + attn_pp_name, + o_proj_pp_name, + intermediate_size, + hidden_act, + mlp_bias, + mlp_pp_name, + norm_eps, + input_norm_pp_name, + post_norm_pp_name, + )?; + layers.push(layers_i); + } + let norm = rms_norm(hidden_size, norm_eps, vb.pp("norm"))?; + let head_dim = head_dim.unwrap_or(hidden_size / num_attention_heads); + let rotary_emb = RoPE::new(head_dim, rope_theta_base, vb.device())?; + Ok(Self { + embed_tokens, + layers, + norm, + rotary_emb, + }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + + let (cos, sin) = self + .rotary_emb + .forward(seqlen_offset, seq_len, inputs_embeds.device())?; + let mut xs = inputs_embeds.clone(); + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + b_size, + seq_len, + 0, + xs.device(), + )?) + } + }; + for layer in self.layers.iter_mut() { + xs = layer.forward(&xs, &cos, &sin, attention_mask.as_ref())?; + } + let xs = xs.apply(&self.norm)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_kv_cache() + } + } +} + +pub struct LlamaForCausalLM { + pub model: LlamaModel, + lm_head: Linear, +} + +impl LlamaForCausalLM { + pub fn new( + vb: VarBuilder, + vocab_size: usize, + hidden_size: usize, + num_hidden_layers: usize, + num_attention_heads: usize, + num_key_value_heads: Option, + head_dim: Option, + attn_bias: bool, + attn_pp_name: &str, + o_proj_pp_name: Option<&str>, + intermediate_size: usize, + hidden_act: Activation, + mlp_bias: bool, + mlp_pp_name: &str, + norm_eps: f64, + input_norm_pp_name: &str, + post_norm_pp_name: &str, + rope_theta_base: f32, + ) -> Result { + let model = LlamaModel::new( + vb.pp("model"), + vocab_size, + hidden_size, + num_hidden_layers, + num_attention_heads, + num_key_value_heads, + head_dim, + attn_bias, + attn_pp_name, + o_proj_pp_name, + intermediate_size, + hidden_act, + mlp_bias, + mlp_pp_name, + norm_eps, + input_norm_pp_name, + post_norm_pp_name, + rope_theta_base, + )?; + let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?; + Ok(Self { model, lm_head }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + let outputs = self.model.forward(inputs_embeds, seqlen_offset)?; + let seq_len = outputs.dim(1)?; + let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&hidden_state)?; + Ok(logits) + } + pub fn clear_kv_cache(&mut self) { + self.model.clear_kv_cache(); + } +} diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs index 6919afb..26764d1 100644 --- a/src/models/deepseek_ocr/model.rs +++ b/src/models/deepseek_ocr/model.rs @@ -271,7 +271,7 @@ impl Block { )?; let norm2 = get_layer_norm(vb.pp("norm2"), eps, dim)?; let mlp_dim = (dim as f32 * mlp_ratio) as usize; - let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, act, true, "lin1", "lin2")?; + let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, dim, act, true, "lin1", "lin2")?; Ok(Self { norm1, attn, diff --git a/src/models/glm_asr_nano/config.rs b/src/models/glm_asr_nano/config.rs new file mode 100644 index 0000000..dc29a7a --- /dev/null +++ b/src/models/glm_asr_nano/config.rs @@ -0,0 +1,88 @@ +use candle_nn::Activation; +use serde::Deserialize; + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrNanoProcessorConfig { + pub audio_token: String, + pub default_transcription_prompt: String, + pub feature_extractor: FeatureExtractor, + pub max_audio_len: usize, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct FeatureExtractor { + pub chunk_length: usize, + pub dither: f32, + pub feature_size: usize, + pub hop_length: usize, + pub n_fft: usize, + pub n_samples: usize, + pub nb_max_frames: usize, + pub padding_side: String, + pub padding_value: f32, + pub return_attention_mask: bool, + pub sampling_rate: usize, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrNanoConfig { + pub audio_config: GlmAsrAudioConfig, + pub audio_token_id: u32, + pub dtype: String, + pub hidden_size: usize, + pub projector_hidden_act: Activation, + pub text_config: GlmAsrTextConfig, + pub vocab_size: usize, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrAudioConfig { + pub attention_dropout: f64, + 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 num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + pub num_mel_bins: usize, + pub partial_rotary_factor: f64, + pub rope_parameters: GlmAsrRopeParameters, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrRopeParameters { + pub partial_rotary_factor: f64, + pub rope_theta: f32, + pub rope_type: String, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrTextConfig { + pub attention_bias: bool, + pub attention_dropout: f64, + pub eos_token_id: Vec, + 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 mlp_bias: bool, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + pub pretraining_tp: usize, + pub rms_norm_eps: f64, + pub rope_parameters: GlmAsrTextRopeParameters, + pub use_cache: bool, + pub vocab_size: usize, +} + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct GlmAsrTextRopeParameters { + pub rope_theta: f32, + pub rope_type: String, +} diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs new file mode 100644 index 0000000..5674531 --- /dev/null +++ b/src/models/glm_asr_nano/generate.rs @@ -0,0 +1,181 @@ +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::{ + chat_template::ChatTemplate, + models::{ + GenerateModel, + glm_asr_nano::{ + config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor, + }, + }, + tokenizer::TokenizerModel, + utils::{ + build_completion_chunk_response, build_completion_response, find_type_files, get_device, + get_dtype, get_logit_processor, + }, +}; + +pub struct GlmAsrNanoGenerateModel<'a> { + chat_template: ChatTemplate<'a>, + tokenizer: TokenizerModel, + processor: GlmAsrNanoProcessor, + glm_asr_nano: GlmAsrNanoModel, + device: Device, + dtype: DType, + eos_token_id1: u32, + eos_token_id2: u32, + eos_token_id3: u32, + model_name: String, +} + +impl<'a> GlmAsrNanoGenerateModel<'a> { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let chat_template = ChatTemplate::init(path)?; + let tokenizer = TokenizerModel::init(path)?; + let device = get_device(device); + let processor = GlmAsrNanoProcessor::new(path, &device, DType::F32)?; + let config_path = path.to_string() + "/config.json"; + let cfg: GlmAsrNanoConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + let cfg_dtype = cfg.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 glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?; + Ok(Self { + chat_template, + tokenizer, + processor, + glm_asr_nano, + device, + dtype, + eos_token_id1: 59246, + eos_token_id2: 59253, + eos_token_id3: 59255, + model_name: "glm-asr-nano".to_string(), + }) + } +} + +impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let seed = match mes.seed { + None => 34562u64, + 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 (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)?; + let mut input_features = Some(input_features.to_dtype(self.dtype)?); + let mut audio_token_lengths = Some(audio_token_lengths); + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let mut generate: Vec = Vec::new(); + let sample_len = mes.max_tokens.unwrap_or(1024); + for _ in 0..sample_len { + let logits = self.glm_asr_nano.forward( + input_features.as_ref(), + audio_token_lengths.as_ref(), + &input_ids, + 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 + || next_token == self.eos_token_id3 + { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + input_features = None; + audio_token_lengths = None; + } + let num_token = generate.len() as u32; + let res = self.tokenizer.token_decode(generate)?; + self.glm_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 seed = match mes.seed { + None => 34562u64, + 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 (input_features, audio_token_lengths, replace_text) = + self.processor.process_info(&mes, &render_text)?; + let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; + + 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 input_features = Some(input_features.to_dtype(self.dtype)?); + let mut audio_token_lengths = Some(audio_token_lengths); + let mut input_ids = input_ids; + for _ in 0..sample_len { + let logits = + self.glm_asr_nano + .forward(input_features.as_ref(), audio_token_lengths.as_ref(), &input_ids, 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)?; + input_features = None; + audio_token_lengths = 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 || next_token == self.eos_token_id3{ + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + input_features = None; + audio_token_lengths = None; + } + self.glm_asr_nano.clear_kv_cache(); + }; + Ok(Box::new(Box::pin(stream))) + } +} diff --git a/src/models/glm_asr_nano/mod.rs b/src/models/glm_asr_nano/mod.rs new file mode 100644 index 0000000..8b1baf7 --- /dev/null +++ b/src/models/glm_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/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs new file mode 100644 index 0000000..21e5023 --- /dev/null +++ b/src/models/glm_asr_nano/model.rs @@ -0,0 +1,320 @@ +use anyhow::Result; +use candle_core::{IndexOp, Tensor}; +use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no_bias}; + +use crate::{ + models::{ + common::{ + LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + }, + glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, + }, + position_embed::rope::{RoPE, glm_asr_apply_rotary_pos_emb}, + utils::tensor_utils::{get_equal_mask, masked_scatter_dim0}, +}; + +#[derive(Debug, Clone)] +// pub struct AttentionNobias { +pub struct GlmAsrAttention { + q_proj: Linear, + k_proj: Linear, + v_proj: Linear, + o_proj: Linear, + num_heads: usize, + num_kv_heads: usize, + num_kv_groups: usize, + head_dim: usize, + middle_size: usize, +} + +impl GlmAsrAttention { + pub fn new( + vb: VarBuilder, + hidden_size: usize, + num_attention_heads: usize, + num_key_value_heads: usize, + head_dim: Option, + ) -> Result { + let num_kv_groups = num_attention_heads / num_key_value_heads; + let head_dim = match head_dim { + None => hidden_size / num_attention_heads, + Some(dim) => dim, + }; + let q_proj = linear(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(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"))?; + + Ok(Self { + q_proj, + k_proj, + v_proj, + o_proj, + num_heads: num_attention_heads, + num_kv_heads: num_key_value_heads, + num_kv_groups, + head_dim, + middle_size: num_attention_heads * head_dim, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: Option<&Tensor>, + sin: Option<&Tensor>, + attention_mask: Option<&Tensor>, + tof32: bool, + ) -> Result { + let (b_sz, q_len, _) = xs.dims3()?; + let query_states = self.q_proj.forward(xs)?; + let key_states = self.k_proj.forward(xs)?; + let value_states = self.v_proj.forward(xs)?; + let query_states = query_states + .reshape((b_sz, q_len, self.num_heads, self.head_dim))? + .transpose(1, 2)?; + let key_states = key_states + .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? + .transpose(1, 2)?; + let value_states = value_states + .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? + .transpose(1, 2)?; + let (query_states, key_states) = if let Some(cos) = cos + && let Some(sin) = sin + { + glm_asr_apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)? + } else { + (query_states, key_states) + }; + + let scale = 1f64 / f64::sqrt(self.head_dim as f64); + let attn_output = eager_attention_forward( + &query_states, + &key_states, + &value_states, + Some(self.num_kv_groups), + attention_mask, + scale, + )?; + let attn_output = attn_output.reshape((b_sz, q_len, self.middle_size))?; + let attn_output = attn_output.apply(&self.o_proj)?; + Ok(attn_output) + } +} + +pub struct GlmAsrEncoderLayer { + self_attn: GlmAsrAttention, + mlp: TwoLinearMLP, + input_layernorm: LayerNorm, + post_attention_layernorm: LayerNorm, +} + +impl GlmAsrEncoderLayer { + pub fn new(vb: VarBuilder, audio_cfg: &GlmAsrAudioConfig) -> Result { + let self_attn = GlmAsrAttention::new( + vb.pp("self_attn"), + audio_cfg.hidden_size, + audio_cfg.num_attention_heads, + audio_cfg.num_key_value_heads, + Some(audio_cfg.head_dim), + )?; + let mlp = TwoLinearMLP::new( + vb.pp("mlp"), + audio_cfg.hidden_size, + audio_cfg.intermediate_size, + audio_cfg.hidden_size, + audio_cfg.hidden_act, + true, + "fc1", + "fc2", + )?; + let input_layernorm = + get_layer_norm(vb.pp("input_layernorm"), 1e-5, audio_cfg.hidden_size)?; + let post_attention_layernorm = get_layer_norm( + vb.pp("post_attention_layernorm"), + 1e-5, + audio_cfg.hidden_size, + )?; + Ok(Self { + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: Option<&Tensor>, + sin: Option<&Tensor>, + attention_mask: Option<&Tensor>, + tof32: bool, + ) -> Result { + let residual = xs.clone(); + let xs = self.input_layernorm.forward(xs)?; + let xs = self + .self_attn + .forward(&xs, cos, sin, attention_mask, tof32)?; + let residual = residual.add(&xs)?; + let xs = self.post_attention_layernorm.forward(&residual)?; + let xs = self.mlp.forward(&xs)?; + let xs = residual.add(&xs)?; + Ok(xs) + } +} + +pub struct GlmAsrEncoder { + conv1: Conv1d, + conv2: Conv1d, + layers: Vec, + norm: LayerNorm, + rotary_emb: RoPE, +} + +impl GlmAsrEncoder { + pub fn new(vb: VarBuilder, audio_cfg: &GlmAsrAudioConfig) -> Result { + let conv1 = get_conv1d( + vb.pp("conv1"), + audio_cfg.num_mel_bins, + audio_cfg.hidden_size, + 3, + 1, + 1, + 1, + 1, + true, + )?; + let conv2 = get_conv1d( + vb.pp("conv2"), + audio_cfg.hidden_size, + audio_cfg.hidden_size, + 3, + 1, + 2, + 1, + 1, + true, + )?; + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + for i in 0..audio_cfg.num_hidden_layers { + let layer_i = GlmAsrEncoderLayer::new(vb_layers.pp(i), audio_cfg)?; + layers.push(layer_i); + } + let norm = get_layer_norm(vb.pp("norm"), 1e-5, audio_cfg.hidden_size)?; + let dim = (audio_cfg.head_dim as f64 * audio_cfg.partial_rotary_factor) as usize; + let rotary_emb = RoPE::new(dim, audio_cfg.rope_parameters.rope_theta, vb.device())?; + Ok(Self { + conv1, + conv2, + layers, + norm, + rotary_emb, + }) + } + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = self.conv1.forward(xs)?.gelu()?; + let xs = self.conv2.forward(&xs)?.gelu()?; + let mut xs = xs.transpose(1, 2)?; + let (_, seq_len, _) = xs.dims3()?; + let (cos, sin) = self.rotary_emb.forward(0, seq_len, xs.device())?; + for encoder_layer in &self.layers { + xs = encoder_layer.forward(&xs, Some(&cos), Some(&sin), None, false)?; + } + let xs = self.norm.forward(&xs)?; + Ok(xs) + } +} + +pub struct GlmAsrNanoModel { + config: GlmAsrNanoConfig, + audio_tower: GlmAsrEncoder, + multi_modal_projector: TwoLinearMLP, + language_model: LlamaForCausalLM, +} + +impl GlmAsrNanoModel { + pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig) -> Result { + let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?; + let multi_modal_projector = TwoLinearMLP::new( + vb.pp("multi_modal_projector"), + config.audio_config.intermediate_size, + config.text_config.hidden_size * 2, + config.text_config.hidden_size, + config.projector_hidden_act, + true, + "linear_1", + "linear_2", + )?; + let language_model = LlamaForCausalLM::new( + vb.pp("language_model"), + config.text_config.vocab_size, + config.text_config.hidden_size, + config.text_config.num_hidden_layers, + config.text_config.num_attention_heads, + Some(config.text_config.num_key_value_heads), + Some(config.text_config.head_dim), + config.text_config.attention_bias, + "self_attn", + Some("o_proj"), + config.text_config.intermediate_size, + config.text_config.hidden_act, + config.text_config.mlp_bias, + "mlp", + config.text_config.rms_norm_eps, + "input_layernorm", + "post_attention_layernorm", + config.text_config.rope_parameters.rope_theta, + )?; + Ok(Self { + config, + audio_tower, + multi_modal_projector, + language_model, + }) + } + + pub fn get_audio_features( + &self, + input_features: &Tensor, + audio_token_lengths: &[u32], + ) -> Result { + let audio_hidden_states = self.audio_tower.forward(input_features)?; + let bs = audio_hidden_states.dim(0)?; + let audio_hidden_states = + audio_hidden_states.reshape((bs, (), self.config.audio_config.intermediate_size))?; + let audio_embeds = self.multi_modal_projector.forward(&audio_hidden_states)?; + let mut valid_audios = vec![]; + for (i, &len) in audio_token_lengths.iter().enumerate() { + let len = len as usize; + let audio_i = audio_embeds.i((i, 0..len, ..))?; + valid_audios.push(audio_i); + } + let audio_embeds = Tensor::cat(&valid_audios, 0)?; + + Ok(audio_embeds) + } + + pub fn forward( + &mut self, + input_features: Option<&Tensor>, + audio_token_lengths: Option<&Vec>, + input_ids: &Tensor, + seqlen_offset: usize, + ) -> Result { + let mut inputs_embeds = self.language_model.model.embed_tokens.forward(input_ids)?; + if let Some(input_features) = input_features + && let Some(audio_token_len) = audio_token_lengths + { + let audio_token_mask = get_equal_mask(input_ids, self.config.audio_token_id)?; + let audio_embeds = self.get_audio_features(input_features, audio_token_len)?; + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embeds, &audio_token_mask)?; + } + let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?; + Ok(logits) + } + pub fn clear_kv_cache(&mut self) { + self.language_model.clear_kv_cache(); + } +} diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs new file mode 100644 index 0000000..043a756 --- /dev/null +++ b/src/models/glm_asr_nano/processor.rs @@ -0,0 +1,235 @@ +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}, + tensor_utils::{pad_reflect_last_dim, split_tensor}, + }, +}; + +pub struct GlmAsrNanoProcessor { + sampling_rate: usize, + chunk_length: usize, + n_samples: usize, + n_fft: usize, + window: Tensor, + mel_filters: Tensor, + hop_length: usize, + audio_token: String, + // audio_token_id: u32, + max_audio_len: usize, + // default_transcription_prompt: String, + device: Device, +} + +impl GlmAsrNanoProcessor { + pub fn new(path: &str, device: &Device, dtype: DType) -> Result { + let path = path.to_string(); + assert!( + std::path::Path::new(&path).exists(), + "model path file not exists" + ); + let processor_config_path = path.to_string() + "/processor_config.json"; + assert!( + std::path::Path::new(&processor_config_path).exists(), + "processor_config.json not exists in model path" + ); + let processor_cfg: GlmAsrNanoProcessorConfig = + serde_json::from_slice(&std::fs::read(processor_config_path)?)?; + let audio_token = processor_cfg.audio_token.clone(); + // let audio_token_id = 59260u32; + let max_audio_len = processor_cfg.max_audio_len; + // let default_transcription_prompt = processor_cfg.default_transcription_prompt.clone(); + let sampling_rate = processor_cfg.feature_extractor.sampling_rate; + let chunk_length = processor_cfg.feature_extractor.chunk_length; + let n_samples = processor_cfg.feature_extractor.n_samples; + let n_fft = processor_cfg.feature_extractor.n_fft; + let hop_length = processor_cfg.feature_extractor.hop_length; + let window = create_hann_window(n_fft, dtype, device)?; + let window = window.unsqueeze(0)?.unsqueeze(0)?; + let mel_filters = mel_filter_bank( + 1 + n_fft / 2, + processor_cfg.feature_extractor.feature_size, + 0.0, + 8000.0, + sampling_rate as f32, + Some("slaney"), + crate::utils::audio_utils::MelScale::Slaney, + false, + device, + )? + .t()?; + Ok(Self { + sampling_rate, + chunk_length, + n_samples, + n_fft, + window, + mel_filters, + hop_length, + audio_token, + // audio_token_id, + max_audio_len, + // default_transcription_prompt, + device: device.clone(), + }) + } + + /// 提取音频帧 + 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 n_frames = (samples - self.n_fft) / self.hop_length + 1; + // (bs, n_frames, n_fft) + let frames = self.extract_frames(&waveform, n_frames)?; + // 应用汉明窗口 + 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 = 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)?; + let ln_spec = mel_spec.log()?; + let log10_spec = ln_spec.broadcast_div(&Tensor::new(f32::ln(10.0), mel_spec.device())?)?; + let max_val = log10_spec.max_all()?.affine(1.0, -8.0)?; + let log10_spec = log10_spec.broadcast_maximum(&max_val)?; + let log_spec = log10_spec.affine(1.0, 4.0)?.affine(1.0 / 4.0, 0.0)?; + Ok(log_spec) + } + + pub fn feature_extractor(&self, raw_speech: Vec) -> Result<(Tensor, Tensor)> { + let mut pad_audio = vec![]; + let mut input_features_mask = vec![]; + for audio in raw_speech { + 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)?; + // (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]); + input_features_mask.push(mask); + } + let input_features = Tensor::cat(&pad_audio, 0)?; + let input_features_mask = Tensor::new(input_features_mask, input_features.device())?; + let input_features = self.extract_fbank_features(&input_features)?; + let (_, audio_len) = input_features_mask.dims2()?; + let mask_idx: Vec = (0..audio_len) + .step_by(self.hop_length) + .map(|i| i as u32) + .collect(); + let mask_idx = Tensor::new(mask_idx, &self.device)?; + let input_features_mask = input_features_mask.index_select(&mask_idx, D::Minus1)?; + Ok((input_features, input_features_mask)) + } + + pub fn process_audio(&self, audios: Vec) -> Result<(Tensor, Tensor, Vec)> { + let window_size = self.sampling_rate * self.chunk_length; + let max_windows = self.max_audio_len / self.chunk_length; + let mut per_sample_windows = vec![]; + let mut flat_chunks = vec![]; + for audio_el in audios { + let audio_el = if audio_el.rank() == 2 { + audio_el.squeeze(0)? + } else { + audio_el + }; + let n_samples = audio_el.dim(0)?; + let n_win = ((n_samples + window_size - 1) / window_size).max(1); + let n_win = if n_win > max_windows { + max_windows + } else { + n_win + }; + per_sample_windows.push(n_win); + let time_cap = (n_win * window_size).min(n_samples); + for i in 0..n_win { + let start = i * window_size; + let end = ((i + 1) * window_size).min(time_cap); + flat_chunks.push(audio_el.i(start..end)?); + } + } + let (input_features, input_features_mask) = self.feature_extractor(flat_chunks)?; + Ok((input_features, input_features_mask, per_sample_windows)) + } + + pub fn get_audio_token_length(&self, audio_lens: Vec) -> Result> { + let merge_factor = 4; + let audio_lens = audio_lens + .iter() + .map(|i| (i + 2 - 3) + 1) // (pad=1, ks=3, stride=1) + .collect::>() + .iter() + .map(|i| (i + 2 - 3) / 2 + 1) // (pad=1, ks=3, stride=2) + .collect::>(); + let num_tokens = audio_lens + .iter() + .map(|i| (i - merge_factor) / merge_factor + 1) + .collect::>(); + Ok(num_tokens) + } + + pub fn process_info( + &self, + mes: &ChatCompletionParameters, + render_text: &str, + ) -> Result<(Tensor, Vec, String)> { + let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?; + let (input_features, input_features_mask, per_sample_windows) = + self.process_audio(audio_tensors)?; + let audio_lengths = input_features_mask.sum(D::Minus1)?; + let audio_vec = split_tensor(&audio_lengths, &per_sample_windows, 0)?; + let audio_vec: Vec = audio_vec + .iter() + .map(|t| t.sum_all().unwrap().to_scalar::().unwrap()) + .collect(); + + let audio_token_lengths = self.get_audio_token_length(audio_vec)?; + let mut text = render_text.to_string(); + for audio_len in audio_token_lengths.clone() { + let replace = "<|placeholder|>".repeat(audio_len as usize); + text = text.replacen(&self.audio_token, &replace, 1); + } + text = text.replace("<|placeholder|>", &self.audio_token); + Ok((input_features, audio_token_lengths, text)) + } +} diff --git a/src/models/mod.rs b/src/models/mod.rs index 705d69e..9e4affb 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -1,5 +1,6 @@ pub mod common; pub mod deepseek_ocr; +pub mod glm_asr_nano; pub mod hunyuan_ocr; pub mod minicpm4; pub mod paddleocr_vl; @@ -16,6 +17,7 @@ use rocket::futures::Stream; use crate::models::{ deepseek_ocr::generate::DeepseekOCRGenerateModel, + 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, @@ -50,6 +52,8 @@ pub enum WhichModel { VoxCPM, #[value(name = "voxcpm1.5")] VoxCPM1_5, + #[value(name = "glm-asr-nano-2512")] + GlmASRNano2512, } pub trait GenerateModel { @@ -76,6 +80,7 @@ pub enum ModelInstance<'a> { PaddleOCRVL(Box>), RMBG2_0(Box), VoxCPM(Box), + GlmASRNano(GlmAsrNanoGenerateModel<'a>), } impl<'a> GenerateModel for ModelInstance<'a> { @@ -89,6 +94,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::PaddleOCRVL(model) => model.generate(mes), ModelInstance::RMBG2_0(model) => model.generate(mes), ModelInstance::VoxCPM(model) => model.generate(mes), + ModelInstance::GlmASRNano(model) => model.generate(mes), } } @@ -112,6 +118,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { ModelInstance::PaddleOCRVL(model) => model.generate_stream(mes), ModelInstance::RMBG2_0(model) => model.generate_stream(mes), ModelInstance::VoxCPM(model) => model.generate_stream(mes), + ModelInstance::GlmASRNano(model) => model.generate_stream(mes), } } } @@ -170,6 +177,10 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result { + let model = GlmAsrNanoGenerateModel::init(path, None, None)?; + ModelInstance::GlmASRNano(model) + } }; Ok(model) } diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index 8872893..f0ec97a 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -205,6 +205,7 @@ impl Qwen3VLVisionBlock { vb.pp("mlp"), config.hidden_size, config.intermediate_size, + config.hidden_size, config.hidden_act, true, "linear_fc1", diff --git a/src/models/rmbg2_0/model.rs b/src/models/rmbg2_0/model.rs index af74a18..9417750 100644 --- a/src/models/rmbg2_0/model.rs +++ b/src/models/rmbg2_0/model.rs @@ -255,7 +255,7 @@ impl SwinTransformerBlock { )?; let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, dim)?; let mlp_dim = (dim as f32 * mlp_ratio) as usize; - let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, act, true, "fc1", "fc2")?; + let mlp = TwoLinearMLP::new(vb.pp("mlp"), dim, mlp_dim, dim, act, true, "fc1", "fc2")?; Ok(Self { norm1, attn, diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 97268b2..a7dee6c 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -213,7 +213,7 @@ impl GenerateModel for VoxCPMGenerate { extract_metadata_value::(&mes.metadata, "retry_badcase_ratio_threshold") .unwrap_or(6.0); let target_text = extract_user_text(&mes)?; - let prompt_wav = extract_audio_url(&mes)?; + let prompt_wav = extract_audio_url(&mes); let prompt_wav_path = if !prompt_wav.is_empty() { Some(prompt_wav[0].clone()) } else { diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs index 08135c9..eaacad3 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -513,7 +513,7 @@ impl VoxCPMModel { let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?; let text_length = text_token.dim(0)?; let mut audio = - load_audio_with_resample(&path, self.device.clone(), Some(self.sample_rate))?; + load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros( @@ -728,11 +728,8 @@ impl VoxCPMModel { ) -> Result> { let text_token = self.tokenizer.encode(prompt_text)?; let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; - let mut audio = load_audio_with_resample( - &prompt_wav_path, - self.device.clone(), - Some(self.sample_rate), - )?; + let mut audio = + load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?; let patch_len = self.patch_size * self.chunk_size; if audio.dim(1)? % patch_len != 0 { audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?; diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index 2189059..8853489 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -115,6 +115,49 @@ pub fn apply_rotary_pos_emb( Ok((q_embed, k_embed)) } +pub fn glm_asr_apply_rotary_pos_emb( + q: &Tensor, + k: &Tensor, + cos: &Tensor, + sin: &Tensor, + tof32: bool, +) -> Result<(Tensor, Tensor)> { + // sin/cos: to (bs, 1, seq_len, head_dim/2) + // q/k: (bs, n_head, seq_len, head_dim) + let mut cos = cos.clone(); + let mut sin = sin.clone(); + if cos.rank() == 2 { + // (seq_len, head_dim/2) -> (1, 1, seq_len, head_dim/2) + cos = cos.unsqueeze(0)?.unsqueeze(0)?; + sin = sin.unsqueeze(0)?.unsqueeze(0)?; + } + if cos.rank() == 3 { + // (bs, seq_len, head_dim/2) -> (bs, 1, seq_len, head_dim/2) + cos = cos.unsqueeze(1)?; + sin = sin.unsqueeze(1)?; + } + let orig_dtype = q.dtype(); + let q = if tof32 { &q.to_dtype(DType::F32)? } else { q }; + let k = if tof32 { &k.to_dtype(DType::F32)? } else { k }; + let cos = cos.to_dtype(q.dtype())?; + let sin = sin.to_dtype(q.dtype())?; + let rotary_dim = cos.dim(D::Minus1)?; + let q_rot = q.narrow(D::Minus1, 0, rotary_dim)?; + let q_pass = q.narrow(D::Minus1, rotary_dim, rotary_dim)?; + let k_rot = k.narrow(D::Minus1, 0, rotary_dim)?; + let k_pass = k.narrow(D::Minus1, rotary_dim, rotary_dim)?; + + let q_embed = q_rot + .broadcast_mul(&cos)? + .add(&rotate_half(&q_rot)?.broadcast_mul(&sin)?)?; + let k_embed = k_rot + .broadcast_mul(&cos)? + .add(&rotate_half(&k_rot)?.broadcast_mul(&sin)?)?; + let q_embed = Tensor::cat(&[q_embed, q_pass], D::Minus1)?.to_dtype(orig_dtype)?; + let k_embed = Tensor::cat(&[k_embed, k_pass], D::Minus1)?.to_dtype(orig_dtype)?; + Ok((q_embed, k_embed)) +} + #[derive(Debug, Clone)] pub struct Qwen2_5VLTextRotaryEmbedding { inv_freq: Vec, diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index 33af770..efb5de1 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -10,12 +10,28 @@ use aha_openai_dive::v1::resources::chat::{ use anyhow::{Result, anyhow}; use base64::Engine; use base64::prelude::BASE64_STANDARD; -use candle_core::{D, Device, Tensor}; +use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_nn::{Conv1d, Conv1dConfig, Module}; +#[cfg(feature = "ffmpeg")] +use ffmpeg_next as ffmpeg; use hound::{SampleFormat, WavReader}; use num::integer::gcd; +use rayon::iter::{IntoParallelRefIterator, ParallelIterator}; +use realfft::RealFftPlanner; +use symphonia::core::audio::{AudioBufferRef, Signal}; +use symphonia::core::codecs::DecoderOptions; +use symphonia::core::formats::FormatOptions; +use symphonia::core::io::MediaSourceStream; +use symphonia::core::meta::MetadataOptions; +use symphonia::core::probe::Hint; +// use rubato::{ +// Async, FixedAsync, Indexing, Resampler, SincInterpolationParameters, SincInterpolationType, +// WindowFunction, +// }; +// use audioadapter_buffers::direct::InterleavedSlice; use crate::utils::get_default_save_dir; +use crate::utils::tensor_utils::linspace; // 重采样方法枚举 #[derive(Debug, Clone, Copy)] @@ -223,6 +239,7 @@ pub fn resample_simple(waveform: &Tensor, orig_freq: i64, new_freq: i64) -> Resu None, ) } + pub fn load_audio_from_url(url: &str) -> Result { tokio::task::block_in_place(|| { let client = reqwest::blocking::Client::new(); @@ -235,7 +252,13 @@ pub fn load_audio_from_url(url: &str) -> Result { } let temp_dir = get_default_save_dir().expect("Failed to get home directory"); let temp_dir = PathBuf::from(temp_dir); - let temp_path = temp_dir.join("temp_audio.wav"); + let temp_path = if url.contains("wav") { + temp_dir.join("temp_audio.wav") + } else if url.contains("mp3") { + temp_dir.join("temp_audio.mp3") + } else { + return Err(anyhow::anyhow!("load audio only surpport wav/mp3 format")); + }; let mut file = std::fs::File::create(&temp_path)?; let mut content = Cursor::new(response.bytes()?); @@ -265,10 +288,19 @@ pub fn get_audio_path(path_str: &str) -> Result { Ok(path) } else if path_str.starts_with("data:audio") && path_str.contains("base64,") { let data: Vec<&str> = path_str.split("base64,").collect(); + let file_mes = data[0]; let data = data[1]; let temp_dir = get_default_save_dir().expect("Failed to get home directory"); let temp_dir = PathBuf::from(temp_dir); - let temp_path = temp_dir.join("temp_audio.wav"); + let temp_path = if file_mes.contains("wav") { + temp_dir.join("temp_audio.wav") + } else if file_mes.contains("mpeg") { + temp_dir.join("temp_audio.mp3") + } else { + return Err(anyhow::anyhow!( + "base64 audio only surpport wav/mpeg(mp3) format" + )); + }; save_audio_from_base64(data, &temp_path)?; Ok(temp_path) } else { @@ -276,8 +308,58 @@ pub fn get_audio_path(path_str: &str) -> Result { } } -pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> { +pub fn load_audio_mono_vec(path: &str) -> Result<(Vec, usize)> { let audio_path = get_audio_path(path)?; + let mut reader = WavReader::open(audio_path)?; + let spec = reader.spec(); + let samples: Vec = match spec.sample_format { + SampleFormat::Int => { + // 将整数样本转换为浮点数 [-1.0, 1.0] + // println!("spec.bits_per_sample: {}", spec.bits_per_sample); + match spec.bits_per_sample { + 8 => reader + .samples::() + .map(|s| s.map(|sample| sample as f32 / i8::MAX as f32)) + .collect::, _>>()?, + 16 => reader + .samples::() + .map(|s| s.map(|sample| sample as f32 / i16::MAX as f32)) + .collect::, _>>()?, + 24 => reader + .samples::() + .map(|s| s.map(|sample| sample as f32 / 8388607.0)) + .collect::, _>>()?, + _ => { + return Err(anyhow::anyhow!( + "Unsupported bit depth: {}", + spec.bits_per_sample + )); + } + } + } + SampleFormat::Float => { + // 直接读取浮点数样本 + reader.samples::().collect::, _>>()? + } + }; + let mono_samples = if spec.channels == 2 { + let mut mono = Vec::with_capacity(samples.len() / 2); + for chunk in samples.chunks(2) { + if chunk.len() == 2 { + mono.push((chunk[0] + chunk[1]) / 2.0); + } + } + mono + } else if spec.channels == 1 { + samples + } else { + return Err(anyhow::anyhow!("only supported mono or stereo")); + }; + let sample_rate = spec.sample_rate as usize; + Ok((mono_samples, sample_rate)) +} + +pub fn load_audio_use_hound(audio_path: PathBuf, device: &Device) -> Result<(Tensor, usize)> { let mut reader = WavReader::open(audio_path)?; let spec = reader.spec(); let samples: Vec = match spec.sample_format { @@ -317,9 +399,10 @@ pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> { samples.len() / spec.channels as usize, spec.channels as usize, ), - &device, + device, )? .t()?; + // println!("audio channels: {}", spec.channels); if spec.channels > 1 { // 对channel通道求平均, channel维度变为1 audio_tensor = audio_tensor.mean_keepdim(0)?; @@ -327,12 +410,109 @@ pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> { Ok((audio_tensor, sample_rate as usize)) } +pub fn load_audio_use_symphonia(path: PathBuf, device: &Device) -> Result<(Tensor, usize)> { + let file = File::open(path.clone())?; + let mss = MediaSourceStream::new(Box::new(file), Default::default()); + + let mut hint = Hint::new(); + let extension = path.extension().and_then(|ext| ext.to_str()).unwrap_or(""); + hint.with_extension(extension); + + let probed = symphonia::default::get_probe().format( + &hint, + mss, + &FormatOptions::default(), + &MetadataOptions::default(), + )?; + + let mut format = probed.format; + let track = format + .default_track() + .ok_or("No default track found") + .map_err(|e| anyhow!("symphonia read err: {}", e))?; + let mut channels = 1; + let sample_rate = track.codec_params.sample_rate.unwrap_or(0); + // 创建解码器 + let mut decoder = + symphonia::default::get_codecs().make(&track.codec_params, &DecoderOptions::default())?; + + // 用于存储所有音频样本的缓冲区 + let mut all_samples: Vec> = Vec::new(); + + // 循环读取数据包并解码 + while let Ok(packet) = format.next_packet() { + match decoder.decode(&packet) { + Ok(decoded) => { + match decoded { + AudioBufferRef::F32(buf) => { + channels = buf.spec().channels.count(); + // 对于浮点格式 + for channel in 0..channels { + if all_samples.len() <= channel { + all_samples.push(Vec::new()); + } + let channel_data = buf.chan(channel); + all_samples[channel].extend_from_slice(channel_data); + } + } + AudioBufferRef::S16(buf) => { + channels = buf.spec().channels.count(); + // 对于16位整数格式,转换为f32 + for channel in 0..channels { + if all_samples.len() <= channel { + all_samples.push(Vec::new()); + } + let channel_data = buf.chan(channel); + let float_samples: Vec = channel_data + .iter() + .map(|&s| s as f32 / 32768.0) // 转换为[-1, 1] + .collect(); + all_samples[channel].extend(float_samples); + } + } + AudioBufferRef::S24(buf) => { + channels = buf.spec().channels.count(); + // 处理24位音频 + for channel in 0..channels { + if all_samples.len() <= channel { + all_samples.push(Vec::new()); + } + let channel_data = buf.chan(channel); + let float_samples: Vec = channel_data + .iter() + .map(|&s| s.inner() as f32 / 8388608.0) // 转换为[-1, 1] + .collect(); + all_samples[channel].extend(float_samples); + } + } + _ => { + println!("不支持的音频格式"); + } + } + } + Err(e) => { + eprintln!("解码错误: {}", e); + break; + } + } + } + let mut audio_tensor = Tensor::new(all_samples, device)?; + if channels > 1 { + // 对channel通道求平均, channel维度变为1 + audio_tensor = audio_tensor.mean_keepdim(0)?; + } + Ok((audio_tensor, sample_rate as usize)) +} + pub fn load_audio_with_resample( path: &str, - device: Device, + device: &Device, target_sample_rate: Option, ) -> Result { - let (mut audio, sr) = load_audio(path, device)?; + let audio_path = get_audio_path(path)?; + // hound 只支持wav文件 + // let (mut audio, sr) = load_audio_use_hound(audio_path, device)?; + let (mut audio, sr) = load_audio_use_symphonia(audio_path, device)?; if let Some(target_sample_rate) = target_sample_rate && target_sample_rate != sr { @@ -387,7 +567,7 @@ pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result> { Ok(wav_buffer) } -pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Result> { +pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Vec { let mut audio_vec = Vec::new(); for chat_mes in mes.messages.clone() { if let ChatMessage::User { content, .. } = chat_mes.clone() @@ -400,20 +580,37 @@ pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Result> } } } - // if let ChatMessage::User { content, .. } = chat_mes.clone() - // && let ChatMessageContent::ContentPart(part_vec) = content - // { - // for part in part_vec { - // if let ChatMessageContentPart::Text(text_part) = part { - // let text = text_part.text; - // if text.chars().count() > 0 { - // ret = ret + &text + "\n" - // } - // } - // } - // } } - Ok(audio_vec) + audio_vec +} + +pub fn extract_audios( + mes: &ChatCompletionParameters, + device: &Device, + target_sample_rate: Option, +) -> Result> { + let audio_url_vec = extract_audio_url(mes); + // 并行加载音频 + audio_url_vec + .par_iter() + .map(|url| load_audio_with_resample(url, device, target_sample_rate)) + .collect() + // #[cfg(not(feature = "ffmpeg"))] + // { + // audio_url_vec + // .par_iter() + // .map(|url| load_audio_with_resample(url, device, target_sample_rate)) + // .collect() + // } + // #[cfg(feature = "ffmpeg")] + // { + // // 该方法wav文件解析有问题 + // use crate::utils::audio_utils::load_and_resample_audio_ffmpeg; + // audio_url_vec + // .par_iter() + // .map(|url| load_and_resample_audio_ffmpeg(url, target_sample_rate, device)) + // .collect() + // } } // 从 ChatCompletionResponse 中提取音频数据 @@ -473,3 +670,369 @@ pub fn extract_and_save_audio_from_response( Ok(saved_files) } + +#[cfg(feature = "ffmpeg")] +pub fn load_and_resample_audio_ffmpeg( + file_path: &str, + target_sample_rate: Option, + device: &Device, +) -> Result { + // 方法只支持mp3 + // wav文件会报错: + // [SWR @ 0x745ff0037840] Input channel layout "" is invalid or unsupported. + // Error: Invalid argument + // 未解决 + ffmpeg::init().map_err(|e| anyhow!(format!("Failed to initialize ffmpeg: {}", e)))?; + + // 打开文件 + let mut ictx = ffmpeg::format::input(&Path::new(file_path)) + .map_err(|e| anyhow!(format!("Failed to open audio file: {}", e)))?; + + // 找到音频流 + let stream = ictx + .streams() + .best(ffmpeg::media::Type::Audio) + .ok_or_else(|| anyhow!(format!("No audio stream found")))?; + let stream_index = stream.index(); + + // 获取解码器 + let codec_params = stream.parameters(); + let mut decoder = ffmpeg::codec::context::Context::from_parameters(codec_params) + .map_err(|e| anyhow!(format!("无法创建解码器上下文: {}", e)))? + .decoder() + .audio() + .map_err(|e| anyhow!(format!("不是音频解码器: {}", e)))?; + + // // 直接更改输入的channel_layout也会报错:Error: Input changed + // let src_channels = decoder.channels(); + // let layout = decoder.channel_layout(); + // if layout.is_empty() || layout.channels() == 0 { + // // 如果没有有效的 channel layout,使用基于通道数的默认布局 + // let layout = ffmpeg::channel_layout::ChannelLayout::default(src_channels as i32); + // decoder.set_channel_layout(layout); + // } + let original_sample_rate = decoder.rate() as usize; + let needs_resampling = match target_sample_rate { + None => false, + Some(target_sr) => target_sr != original_sample_rate, + }; + // 存储音频数据 + let mut audio_buffer = vec![]; + if !needs_resampling { + // 不需要重采样,直接解码音频 + for (stream, packet) in ictx.packets() { + if stream.index() == stream_index { + decoder.send_packet(&packet)?; + let mut decoded = ffmpeg::util::frame::Audio::empty(); + while decoder.receive_frame(&mut decoded).is_ok() { + let planes = decoded.planes(); + if planes == 1 { + let data_slice = decoded.plane::(0); + audio_buffer.extend_from_slice(data_slice); + } else { + let mut channel_data: Vec<&[f32]> = vec![]; + for plane_idx in 0..planes { + let plane_data = decoded.plane::(plane_idx); + channel_data.push(plane_data); + } + let channel_len = channel_data[0].len(); + for sample_idx in 0..channel_len { + let mut sum = 0.0f32; + for channel in &channel_data { + sum += channel[sample_idx]; + } + let avg = sum / planes as f32; + audio_buffer.push(avg); + } + } + } + } + } + } else { + let target_sample_rate = target_sample_rate.unwrap_or(16000); + // 创建重采样器, 通道为1 + let mut resampler = ffmpeg::software::resampling::context::Context::get( + decoder.format(), + decoder.channel_layout(), + decoder.rate() as u32, + ffmpeg::format::Sample::F32(ffmpeg::format::sample::Type::Planar), + ffmpeg::channel_layout::ChannelLayout::default(1), + target_sample_rate as u32, + ) + .map_err(|e| anyhow!(format!("无法创建重采样器: {}", e)))?; + + // let mut resampler = decoder.resampler( + // ffmpeg::format::Sample::F32(ffmpeg::format::sample::Type::Planar), + // ffmpeg::channel_layout::ChannelLayout::default(target_channels as i32), + // target_sample_rate, + // )?; + + // 处理所有包 + for (stream, packet) in ictx.packets() { + if stream.index() == stream_index { + // 解码 + decoder.send_packet(&packet)?; + + let mut decoded = ffmpeg::util::frame::Audio::empty(); + while decoder.receive_frame(&mut decoded).is_ok() { + // 重采样 + let mut resampled = ffmpeg::util::frame::Audio::empty(); + resampler.run(&decoded, &mut resampled)?; + + // 提取数据,Planar格式 + let data_slice = resampled.plane::(0); + audio_buffer.extend_from_slice(data_slice); + } + } + } + + // 处理剩余数据 + decoder.send_eof()?; + + let mut decoded = ffmpeg::util::frame::Audio::empty(); + while decoder.receive_frame(&mut decoded).is_ok() { + let mut resampled = ffmpeg::util::frame::Audio::empty(); + resampler.run(&decoded, &mut resampled)?; + + // 提取数据,Planar格式 + let data_slice = resampled.plane::(0); + audio_buffer.extend_from_slice(data_slice); + } + } + + let audio_tensor = Tensor::new(audio_buffer, device)?; + + Ok(audio_tensor) +} + +// pub fn load_and_resample_audio_rubato( +// file_path: &str, +// target_sample_rate: usize, +// device: &Device, +// ) -> Result { +// let (mono_audio, ori_sample_rate) = load_audio_mono_vec(file_path)?; +// let params = SincInterpolationParameters { +// sinc_len: 256, +// f_cutoff: 0.95, +// interpolation: SincInterpolationType::Cubic, +// oversampling_factor: 256, +// window: WindowFunction::BlackmanHarris2, +// }; +// let input_len = mono_audio.len(); +// let mut resampler = Async::::new_sinc( +// target_sample_rate as f64 / ori_sample_rate as f64, // 重采样比例 +// 1.0, // 输出/输入采样率比 +// ¶ms, +// input_len, +// 1, // 单通道 +// FixedAsync::Input, +// ) +// .map_err(|e| anyhow!(format!("无法创建重采样器: {}", e)))?; + +// let mono_audio: Vec = mono_audio.iter().map(|x| *x as f64).collect(); +// let input_adapter = InterleavedSlice::new(&mono_audio, 1, input_len)?; + +// let mut outdata = vec![0.0f64; input_len * 2]; +// let mut output_adapter = InterleavedSlice::new_mut(&mut outdata, 1, input_len * 2)?; +// // Preparations +// let mut indexing = Indexing { +// input_offset: 0, +// output_offset: 0, +// active_channels_mask: None, +// partial_len: None, +// }; +// let mut input_frames_left = input_len; +// let mut input_frames_next = resampler.input_frames_max(); +// while input_frames_left >= input_frames_next { +// let (frames_read, frames_written) = +// resampler.process_into_buffer(&input_adapter, &mut output_adapter, Some(&indexing))?; +// indexing.input_offset += frames_read; +// indexing.output_offset += frames_written; +// input_frames_left -= frames_read; +// input_frames_next = resampler.input_frames_next(); +// } +// indexing.partial_len = Some(input_frames_left); +// let (_nbr_in, _nbr_out) = resampler +// .process_into_buffer(&input_adapter, &mut output_adapter, Some(&indexing)) +// .unwrap(); +// let output_len = input_len * target_sample_rate / ori_sample_rate; +// let audio_tensor = +// Tensor::new(&outdata[0..output_len], device)?.to_dtype(candle_core::DType::F32)?; +// Ok(audio_tensor) +// } + +pub fn create_hann_window(window_size: usize, dtype: DType, device: &Device) -> Result { + let n = window_size as f64; + let window: Vec = (0..window_size) + .map(|i| { + let i_f64 = i as f64; + let val = 0.5 * (1.0 - (2.0 * PI * i_f64 / n).cos()); + val as f32 + }) + .collect(); + Ok(Tensor::from_vec(window, window_size, device)?.to_dtype(dtype)?) +} + +/// 梅尔频率刻度类型 +#[derive(Debug, Clone, Copy)] +pub enum MelScale { + Htk, + Kaldi, + Slaney, +} + +/// 将赫兹转换为梅尔频率 +pub fn hertz_to_mel(freq: f32, mel_scale: MelScale) -> f32 { + match mel_scale { + MelScale::Htk => 2595.0 * ((1.0 + freq / 700.0).log10()), + MelScale::Kaldi => 1127.0 * ((1.0 + freq / 700.0).ln()), + MelScale::Slaney => { + let min_log_hertz = 1000.0; + let min_log_mel = 15.0; + let logstep = 27.0 / 6.4_f32.ln(); + let mut mels = 3.0 * freq / 200.0; + + if freq >= min_log_hertz { + mels = min_log_mel + (freq / min_log_hertz).ln() * logstep; + } + mels + } + } +} + +/// 将梅尔频率转换为赫兹 +pub fn mel_to_hertz(mels: f32, mel_scale: MelScale) -> f32 { + match mel_scale { + MelScale::Htk => 700.0 * (10.0_f32.powf(mels / 2595.0) - 1.0), + MelScale::Kaldi => 700.0 * (f32::exp(mels / 1127.0) - 1.0), + MelScale::Slaney => { + let min_log_hertz = 1000.0; + let min_log_mel = 15.0; + let logstep = 6.4_f32.ln() / 27.0; + let mut freq = 200.0 * mels / 3.0; + + if mels >= min_log_mel { + freq = min_log_hertz * f32::exp(logstep * (mels - min_log_mel)); + } + freq + } + } +} + +pub fn create_triangular_filter_bank(fft_freqs: &Tensor, filter_freqs: &Tensor) -> Result { + // fft_freqs/filter_freqs -> 1d + let len = filter_freqs.dim(0)?; + let filter_diff = filter_freqs + .narrow(0, 1, len - 1)? + .sub(&filter_freqs.narrow(0, 0, len - 1)?)?; + let slopes = filter_freqs + .unsqueeze(0)? + .broadcast_sub(&fft_freqs.unsqueeze(1)?)?; + let down_slopes = slopes + .narrow(D::Minus1, 0, len - 2)? + .affine(-1.0, 0.0)? + .broadcast_div(&filter_diff.narrow(0, 0, len - 2)?)?; + let up_slopes = slopes + .narrow(D::Minus1, 2, len - 2)? + .broadcast_div(&filter_diff.narrow(0, 1, len - 2)?)?; + let res = down_slopes + .minimum(&up_slopes)? + .maximum(&Tensor::zeros_like(&down_slopes)?)?; + Ok(res) +} + +/// 创建梅尔滤波器组 +pub fn mel_filter_bank( + num_frequency_bins: usize, + num_mel_filters: usize, + min_frequency: f32, + max_frequency: f32, + sampling_rate: f32, + norm: Option<&str>, + mel_scale: MelScale, + triangularize_in_mel_space: bool, + device: &Device, +) -> Result { + // 参数验证 + if let Some(n) = norm + && n != "slaney" + { + return Err(anyhow::anyhow!("norm must be one of None or 'slaney'")); + } + if num_frequency_bins < 2 { + return Err(anyhow::anyhow!( + "Require num_frequency_bins: {} >= 2", + num_frequency_bins + )); + } + if min_frequency > max_frequency { + return Err(anyhow::anyhow!( + "Require min_frequency: {} <= max_frequency: {}", + min_frequency, + max_frequency + )); + } + // 计算梅尔频率范围 + let mel_min = hertz_to_mel(min_frequency, mel_scale); + let mel_max = hertz_to_mel(max_frequency, mel_scale); + + // 在梅尔刻度上均匀分布频率点(包括边界点) + let mel_freqs = linspace(mel_min, mel_max, num_mel_filters + 2, device)?; + + // 将梅尔频率转换回赫兹频率 + let filter_freqs: Vec = mel_freqs + .to_vec1::()? + .iter() + .map(|&m| mel_to_hertz(m, mel_scale)) + .collect(); + let mut filter_freqs = Tensor::new(filter_freqs, device)?; + + let fft_freqs = if triangularize_in_mel_space { + // 在梅尔空间中应用三角滤波器 + let fft_bin_width = sampling_rate / ((num_frequency_bins as f32 - 1.0) * 2.0); + let fft_vec: Vec = (0..num_frequency_bins) + .map(|i| hertz_to_mel(fft_bin_width * i as f32, mel_scale)) + .collect(); + filter_freqs = mel_freqs; + Tensor::new(fft_vec, device)? + } else { + // 在赫兹频率上 + linspace(0.0, sampling_rate / 2.0, num_frequency_bins, device)? + }; + + // 创建三角滤波器组 + let mut mel_filters = create_triangular_filter_bank(&fft_freqs, &filter_freqs)?; + + // 如果需要,进行归一化 + if let Some(n) = norm + && n == "slaney" + { + // Slaney风格的归一化 + let enorm = (2.0 + / filter_freqs + .i(2..num_mel_filters + 2)? + .sub(&filter_freqs.i(0..num_mel_filters)?)?)? + .unsqueeze(0)?; + mel_filters = mel_filters.broadcast_mul(&enorm)?; + } + + // // 检查是否有零值滤波器 + // let mel_max = mel_filters.max(0)?; + // let mel_max_eq_zero = mel_max.eq(&Tensor::zeros_like(&mel_max)?)?; + // let eq_zero_index = zero_index_vec(&mel_max_eq_zero)?; + // if eq_zero_index.len() > 0 { + // println!("At least one mel filter has all zero values."); + // } + + Ok(mel_filters) +} + +pub fn stft_audio(n_fft: usize, frame_wave: &[f32]) -> Result> { + let mut real_planner = RealFftPlanner::::new(); + let r2c = real_planner.plan_fft_forward(n_fft); + let mut spectrum = r2c.make_output_vec(); + let mut frame_wave = frame_wave.to_owned(); + r2c.process(&mut frame_wave, &mut spectrum)?; + let output: Vec = spectrum.iter().map(|complex| complex.norm_sqr()).collect(); + Ok(output) +} diff --git a/src/utils/img_utils.rs b/src/utils/img_utils.rs index 7120ab0..78838e4 100644 --- a/src/utils/img_utils.rs +++ b/src/utils/img_utils.rs @@ -8,7 +8,7 @@ use anyhow::{Result, anyhow}; use base64::{Engine, engine::general_purpose}; use candle_core::{DType, Device, Tensor}; use image::{DynamicImage, ImageBuffer, ImageReader, Rgb, RgbImage, imageops}; -use rayon::prelude::*; +use rayon::iter::{IntoParallelRefIterator, ParallelIterator}; use crate::utils::{ceil_by_factor, floor_by_factor, round_by_factor}; diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index fbee88d..2ee4a2b 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -844,3 +844,29 @@ pub fn nonzero(input: &Tensor) -> Result<(Vec, Vec)> { } Ok((topk_ids, token_ids_all)) } + +pub fn pad_reflect_last_dim(t: &Tensor, pad: (usize, usize)) -> Result { + let (pad_l, pad_r) = pad; + let last_dim = t.dim(D::Minus1)?; + if pad_l >= last_dim || pad_r >= last_dim { + return Err(anyhow!(format!( + "input pad_l {}, pad_r {} must less than t last_dim: {}", + pad_l, pad_r, last_dim + ))); + } + let mut pad_tensor = t.clone(); + if pad_l > 0 { + let left = pad_tensor.narrow(D::Minus1, 1, pad_l)?.contiguous()?; + let last_dim_id = left.rank() - 1; + let left_flip = left.flip(&[last_dim_id])?; + pad_tensor = Tensor::cat(&[&left_flip, &pad_tensor], D::Minus1)?; + } + if pad_r > 0 { + let start_i = last_dim - pad_r; + let right = pad_tensor.narrow(D::Minus1, start_i, pad_r)?.contiguous()?; + let last_dim_id = right.rank() - 1; + let right_flip = right.flip(&[last_dim_id])?; + pad_tensor = Tensor::cat(&[&pad_tensor, &right_flip], D::Minus1)?; + } + Ok(pad_tensor) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 11558d9..e0d083f 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,14 +1,31 @@ +use aha::utils::audio_utils::create_hann_window; use anyhow::Result; -use candle_core::Tensor; +use candle_core::DType; #[test] fn messy_test() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg messy_test -r -- --nocapture let device = &candle_core::Device::Cpu; - // let path = get_default_save_dir(); - let x = Tensor::arange(0.0, 9.0, device)?; - println!("x: {}", x); + let window = create_hann_window(400, DType::F32, device)?; + println!("window: {}", window); + // let audio_path = "file:///home/jhq/Videos/voice_01.wav"; + // let audio_path = "/home/jhq/Videos/zh.mp3"; + // let audio_path = "/home/jhq/Videos/zh.mp3"; + // // let audio_tensor = load_and_resample_audio_rubato(audio_path, 16000, device)?; + // // let audio_tensor = load_audio_with_resample(audio_path, device, Some(16000))?; + // // println!("audio_tensor: {}", audio_tensor); + // #[cfg(feature = "ffmpeg")] + // { + // use aha::utils::audio_utils::load_and_resample_audio_ffmpeg; + // let audio_tensor = load_and_resample_audio_ffmpeg(audio_path, Some(16000), device)?; + // println!("audio_tensor: {}", audio_tensor); + // } + + // // let path = get_default_save_dir(); + // // let x = Tensor::new(array, device) + // let x = Tensor::arange(0.0, 9.0, device)?; + // println!("x: {}", x); // let x = x // .unsqueeze(0)? // .unsqueeze(0)? diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs new file mode 100644 index 0000000..7414cb5 --- /dev/null +++ b/tests/test_glm_asr_nano.rs @@ -0,0 +1,94 @@ +use std::{pin::pin, time::Instant}; + +use aha::models::{GenerateModel, glm_asr_nano::generate::GlmAsrNanoGenerateModel}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use rocket::futures::StreamExt; + +#[test] +fn glm_asr_nano_generate() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_generate -r -- --nocapture + let model_path = "/home/jhq/huggingface_model/ZhipuAI/GLM-ASR-Nano-2512/"; + let message = r#" + { + "model": "glm-asr-nano", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.mp3" + } + }, + { + "type": "text", + "text": "Please transcribe this audio into text" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut glm_asr_model = GlmAsrNanoGenerateModel::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 = glm_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 glm_asr_nano_stream() -> Result<()> { + // RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_stream -r -- --nocapture + let model_path = "/home/jhq/huggingface_model/ZhipuAI/GLM-ASR-Nano-2512/"; + let message = r#" + { + "model": "glm-asr-nano", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "audio", + "audio_url": + { + "url": "file://./assets/audio/zh.mp3" + } + }, + { + "type": "text", + "text": "Please transcribe this audio into text" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut glm_asr_model = GlmAsrNanoGenerateModel::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!(glm_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/weight_test.rs b/tests/weight_test.rs index ab68db6..bbfa740 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -119,3 +119,22 @@ fn hunyuanocr_weight() -> Result<()> { println!("model_list: {:?}", model_list); Ok(()) } + +#[test] +fn glm_asr_nano_weight() -> Result<()> { + let model_path = "/home/jhq/huggingface_model/ZhipuAI/GLM-ASR-Nano-2512/"; + 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(()) +}