Merge branch 'add_model'

This commit is contained in:
jhqxxx
2026-01-07 21:48:45 +08:00
24 changed files with 2062 additions and 58 deletions
Generated
+211 -9
View File
@@ -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]]
+7 -2
View File
@@ -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"
+7 -1
View File
@@ -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 模型
Binary file not shown.
+1
View File
@@ -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(),
+201 -10
View File
@@ -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<Self> {
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<Conv1d> {
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<LayerNorm> {
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<NaiveAttnGateUpDownMLPBlock>,
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<usize>,
head_dim: Option<usize>,
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<Self> {
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<Tensor> {
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<Tensor> = {
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<usize>,
head_dim: Option<usize>,
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<Self> {
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<Tensor> {
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();
}
}
+1 -1
View File
@@ -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,
+88
View File
@@ -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<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 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,
}
+181
View File
@@ -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<DType>) -> Result<Self> {
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<ChatCompletionResponse> {
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<u32> = 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<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
}
}
+4
View File
@@ -0,0 +1,4 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
+320
View File
@@ -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<usize>,
) -> Result<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<GlmAsrEncoderLayer>,
norm: LayerNorm,
rotary_emb: RoPE,
}
impl GlmAsrEncoder {
pub fn new(vb: VarBuilder, audio_cfg: &GlmAsrAudioConfig) -> Result<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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<u32>>,
input_ids: &Tensor,
seqlen_offset: usize,
) -> Result<Tensor> {
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();
}
}
+235
View File
@@ -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<Self> {
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<Tensor> {
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<Tensor> {
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::<f32>()?;
let wave_i_fft_vec: Result<Vec<Vec<f32>>> = 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<Tensor>) -> 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<u32> = (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<Tensor>) -> Result<(Tensor, Tensor, Vec<usize>)> {
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<u32>) -> Result<Vec<u32>> {
let merge_factor = 4;
let audio_lens = audio_lens
.iter()
.map(|i| (i + 2 - 3) + 1) // (pad=1, ks=3, stride=1)
.collect::<Vec<u32>>()
.iter()
.map(|i| (i + 2 - 3) / 2 + 1) // (pad=1, ks=3, stride=2)
.collect::<Vec<u32>>();
let num_tokens = audio_lens
.iter()
.map(|i| (i - merge_factor) / merge_factor + 1)
.collect::<Vec<u32>>();
Ok(num_tokens)
}
pub fn process_info(
&self,
mes: &ChatCompletionParameters,
render_text: &str,
) -> Result<(Tensor, Vec<u32>, 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<u32> = audio_vec
.iter()
.map(|t| t.sum_all().unwrap().to_scalar::<u32>().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))
}
}
+11
View File
@@ -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<PaddleOCRVLGenerateModel<'a>>),
RMBG2_0(Box<RMBG2_0Model>),
VoxCPM(Box<VoxCPMGenerate>),
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<ModelInstance<'_
let model = VoxCPMGenerate::init(path, None, None)?;
ModelInstance::VoxCPM(Box::new(model))
}
WhichModel::GlmASRNano2512 => {
let model = GlmAsrNanoGenerateModel::init(path, None, None)?;
ModelInstance::GlmASRNano(model)
}
};
Ok(model)
}
+1
View File
@@ -205,6 +205,7 @@ impl Qwen3VLVisionBlock {
vb.pp("mlp"),
config.hidden_size,
config.intermediate_size,
config.hidden_size,
config.hidden_act,
true,
"linear_fc1",
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -213,7 +213,7 @@ impl GenerateModel for VoxCPMGenerate {
extract_metadata_value::<f64>(&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 {
+3 -6
View File
@@ -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<HashMap<String, Tensor>> {
let text_token = self.tokenizer.encode(prompt_text)?;
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
let mut audio = load_audio_with_resample(
&prompt_wav_path,
self.device.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)?;
+43
View File
@@ -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<f32>,
+584 -21
View File
@@ -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<PathBuf> {
tokio::task::block_in_place(|| {
let client = reqwest::blocking::Client::new();
@@ -235,7 +252,13 @@ pub fn load_audio_from_url(url: &str) -> Result<PathBuf> {
}
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<PathBuf> {
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<PathBuf> {
}
}
pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> {
pub fn load_audio_mono_vec(path: &str) -> Result<(Vec<f32>, usize)> {
let audio_path = get_audio_path(path)?;
let mut reader = WavReader::open(audio_path)?;
let spec = reader.spec();
let samples: Vec<f32> = 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::<i8>()
.map(|s| s.map(|sample| sample as f32 / i8::MAX as f32))
.collect::<Result<Vec<_>, _>>()?,
16 => reader
.samples::<i16>()
.map(|s| s.map(|sample| sample as f32 / i16::MAX as f32))
.collect::<Result<Vec<_>, _>>()?,
24 => reader
.samples::<i32>()
.map(|s| s.map(|sample| sample as f32 / 8388607.0))
.collect::<Result<Vec<_>, _>>()?,
_ => {
return Err(anyhow::anyhow!(
"Unsupported bit depth: {}",
spec.bits_per_sample
));
}
}
}
SampleFormat::Float => {
// 直接读取浮点数样本
reader.samples::<f32>().collect::<Result<Vec<_>, _>>()?
}
};
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<f32> = 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<f32>> = 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<f32> = 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<f32> = 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<usize>,
) -> Result<Tensor> {
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<Vec<u8>> {
Ok(wav_buffer)
}
pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Result<Vec<String>> {
pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Vec<String> {
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<Vec<String>>
}
}
}
// 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<usize>,
) -> Result<Vec<Tensor>> {
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<usize>,
device: &Device,
) -> Result<Tensor> {
// 方法只支持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::<f32>(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::<f32>(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::<f32>(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::<f32>(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<Tensor> {
// 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::<f64>::new_sinc(
// target_sample_rate as f64 / ori_sample_rate as f64, // 重采样比例
// 1.0, // 输出/输入采样率比
// &params,
// input_len,
// 1, // 单通道
// FixedAsync::Input,
// )
// .map_err(|e| anyhow!(format!("无法创建重采样器: {}", e)))?;
// let mono_audio: Vec<f64> = 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<Tensor> {
let n = window_size as f64;
let window: Vec<f32> = (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<Tensor> {
// 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<Tensor> {
// 参数验证
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<f32> = mel_freqs
.to_vec1::<f32>()?
.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<f32> = (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<Vec<f32>> {
let mut real_planner = RealFftPlanner::<f32>::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<f32> = spectrum.iter().map(|complex| complex.norm_sqr()).collect();
Ok(output)
}
+1 -1
View File
@@ -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};
+26
View File
@@ -844,3 +844,29 @@ pub fn nonzero(input: &Tensor) -> Result<(Vec<u32>, Vec<u32>)> {
}
Ok((topk_ids, token_ids_all))
}
pub fn pad_reflect_last_dim(t: &Tensor, pad: (usize, usize)) -> Result<Tensor> {
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)
}
+22 -5
View File
@@ -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)?
+94
View File
@@ -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(())
}
+19
View File
@@ -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(())
}