add glm-asr-nano
This commit is contained in:
Generated
+211
-9
@@ -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
@@ -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"
|
||||
|
||||
@@ -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.
@@ -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
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use serde::{Deserialize};
|
||||
use candle_nn::Activation;
|
||||
use serde::Deserialize;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Deserialize)]
|
||||
pub struct GlmAsrNanoProcessorConfig {
|
||||
@@ -21,4 +22,67 @@ pub struct FeatureExtractor {
|
||||
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,
|
||||
}
|
||||
|
||||
@@ -1,24 +1,37 @@
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
use candle_core::{DType, Device};
|
||||
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::glm_asr_nano::{config::GlmAsrNanoProcessorConfig, processor::GlmAsrNanoProcessor},
|
||||
models::{
|
||||
GenerateModel,
|
||||
glm_asr_nano::{
|
||||
config::GlmAsrNanoConfig, model::GlmAsrNanoModel, processor::GlmAsrNanoProcessor,
|
||||
},
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{get_device, get_dtype},
|
||||
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,
|
||||
glm_asr_nano: GlmAsrNanoModel,
|
||||
device: Device,
|
||||
// eos_token_id1: u32,
|
||||
// eos_token_id2: u32,
|
||||
// eos_token_id3: u32,
|
||||
// generation_config: GlmAsrNanoGenerationConfig,
|
||||
dtype: DType,
|
||||
eos_token_id1: u32,
|
||||
eos_token_id2: u32,
|
||||
eos_token_id3: u32,
|
||||
model_name: String,
|
||||
}
|
||||
|
||||
@@ -28,23 +41,141 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
|
||||
let tokenizer = TokenizerModel::init(path)?;
|
||||
let device = get_device(device);
|
||||
let processor = GlmAsrNanoProcessor::new(path, &device, DType::F32)?;
|
||||
// let cfg_dtype = cfg.dtype.as_str();
|
||||
// let dtype = get_dtype(dtype, cfg_dtype);
|
||||
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,
|
||||
// eos_token_id1,
|
||||
// eos_token_id2,
|
||||
// eos_token_id3,
|
||||
dtype,
|
||||
eos_token_id1: 59246,
|
||||
eos_token_id2: 59253,
|
||||
eos_token_id3: 59255,
|
||||
model_name: "glm-asr-nano".to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn generate(&self, mes: ChatCompletionParameters) -> Result<()> {
|
||||
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 audio = self.processor.process_info(&mes)?;
|
||||
Ok(())
|
||||
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)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
pub mod processor;
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,30 +1,30 @@
|
||||
use std::f32;
|
||||
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
use candle_core::{DType, Device, IndexOp, Tensor};
|
||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
|
||||
|
||||
use crate::{
|
||||
models::glm_asr_nano::config::GlmAsrNanoProcessorConfig,
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{audio_utils::extract_audios, extract_user_text},
|
||||
utils::{
|
||||
audio_utils::{create_hann_window, extract_audios, mel_filter_bank, stft_audio},
|
||||
tensor_utils::{pad_reflect_last_dim, split_tensor},
|
||||
},
|
||||
};
|
||||
|
||||
pub struct WhisperFeatureExtractor {
|
||||
feature_size: usize,
|
||||
sampling_rate: usize,
|
||||
padding_value: f32,
|
||||
hop_length: usize,
|
||||
chunk_length: usize,
|
||||
n_fft: usize,
|
||||
dither: f32,
|
||||
}
|
||||
|
||||
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,
|
||||
// audio_token_id: u32,
|
||||
max_audio_len: usize,
|
||||
default_transcription_prompt: String,
|
||||
// default_transcription_prompt: String,
|
||||
device: Device,
|
||||
}
|
||||
|
||||
@@ -43,52 +43,193 @@ impl GlmAsrNanoProcessor {
|
||||
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 audio_token_id = 59260u32;
|
||||
let max_audio_len = processor_cfg.max_audio_len;
|
||||
let default_transcription_prompt = processor_cfg.default_transcription_prompt.clone();
|
||||
// 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,
|
||||
// audio_token_id,
|
||||
max_audio_len,
|
||||
default_transcription_prompt,
|
||||
// default_transcription_prompt,
|
||||
device: device.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
// pub fn process_audio(&self, audios: Vec<Tensor>) -> Result<Tensor> {
|
||||
// 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 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)?);
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
/// 提取音频帧
|
||||
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
|
||||
) -> Result<Tensor> {
|
||||
mes: &ChatCompletionParameters,
|
||||
render_text: &str,
|
||||
) -> Result<(Tensor, Vec<u32>, String)> {
|
||||
let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?;
|
||||
println!("audio: {}", audio_tensors[0]);
|
||||
// let audio = self.process_audio(audio_tensors)?;
|
||||
Ok(audio_tensors[0].clone())
|
||||
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
-1
@@ -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;
|
||||
@@ -7,7 +8,6 @@ pub mod qwen2_5vl;
|
||||
pub mod qwen3vl;
|
||||
pub mod rmbg2_0;
|
||||
pub mod voxcpm;
|
||||
pub mod glm_asr_nano;
|
||||
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
@@ -17,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,
|
||||
@@ -51,6 +52,8 @@ pub enum WhichModel {
|
||||
VoxCPM,
|
||||
#[value(name = "voxcpm1.5")]
|
||||
VoxCPM1_5,
|
||||
#[value(name = "glm-asr-nano-2512")]
|
||||
GlmASRNano2512,
|
||||
}
|
||||
|
||||
pub trait GenerateModel {
|
||||
@@ -77,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> {
|
||||
@@ -90,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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -113,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),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -171,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)
|
||||
}
|
||||
|
||||
@@ -205,6 +205,7 @@ impl Qwen3VLVisionBlock {
|
||||
vb.pp("mlp"),
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
config.hidden_size,
|
||||
config.hidden_act,
|
||||
true,
|
||||
"linear_fc1",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
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)?;
|
||||
|
||||
@@ -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>,
|
||||
|
||||
+575
-7
@@ -10,13 +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)]
|
||||
@@ -224,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();
|
||||
@@ -236,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()?);
|
||||
@@ -266,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 {
|
||||
@@ -277,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 {
|
||||
@@ -321,6 +402,7 @@ pub fn load_audio(path: &str, device: &Device) -> Result<(Tensor, usize)> {
|
||||
device,
|
||||
)?
|
||||
.t()?;
|
||||
// println!("audio channels: {}", spec.channels);
|
||||
if spec.channels > 1 {
|
||||
// 对channel通道求平均, channel维度变为1
|
||||
audio_tensor = audio_tensor.mean_keepdim(0)?;
|
||||
@@ -328,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,
|
||||
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
|
||||
{
|
||||
@@ -405,10 +584,33 @@ pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Vec<String> {
|
||||
audio_vec
|
||||
}
|
||||
|
||||
pub fn extract_audios(mes: &ChatCompletionParameters, device: &Device, target_sample_rate: Option<usize>) -> Result<Vec<Tensor>> {
|
||||
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()
|
||||
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 中提取音频数据
|
||||
@@ -468,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, // 输出/输入采样率比
|
||||
// ¶ms,
|
||||
// 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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
+23
-5
@@ -1,13 +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)?
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
use aha::models::glm_asr_nano::generate::GlmAsrNanoGenerateModel;
|
||||
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 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/zai-org/GLM-ASR-Nano-2512/";
|
||||
let model_path = "/home/jhq/huggingface_model/ZhipuAI/GLM-ASR-Nano-2512/";
|
||||
let message = r#"
|
||||
{
|
||||
"model": "glm-asr-nano",
|
||||
@@ -17,9 +20,9 @@ fn glm_asr_nano_generate() -> Result<()> {
|
||||
"type": "audio",
|
||||
"audio_url":
|
||||
{
|
||||
"url": "file://./assets/audio/voice_01.wav"
|
||||
"url": "https://sis-sample-audio.obs.cn-north-1.myhuaweicloud.com/16k16bit.mp3"
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Please transcribe this audio into text"
|
||||
@@ -30,7 +33,62 @@ fn glm_asr_nano_generate() -> Result<()> {
|
||||
}
|
||||
"#;
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
let glm_asr_model = GlmAsrNanoGenerateModel::init(model_path, None, None)?;
|
||||
let _ = glm_asr_model.generate(mes)?;
|
||||
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(())
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user