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