From adf31dc16fcb853f6f95de9ecff526e841cde7f0 Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Wed, 8 Apr 2026 18:55:03 +0800
Subject: [PATCH] add VoxCPM2
---
Cargo.toml | 2 +-
README.md | 5 +-
README.zh-CN.md | 5 +-
docs/changelog.md | 3 +
docs/changelog.zh-CN.md | 3 +
docs/model-card.zh-CN.md | 3 +
docs/supported-models.md | 2 +-
docs/supported-models.zh-CN.md | 2 +-
scripts/download_and_run.sh | 4 +
src/cli/mod.rs | 2 +-
src/models/common/model_mapping.rs | 4 +-
src/models/common/modules.rs | 21 ++-
src/models/deepseek_ocr/model.rs | 6 +-
src/models/minicpm4/model.rs | 6 +-
src/models/mod.rs | 2 +-
src/models/voxcpm/audio_vae.rs | 132 +++++++++++++-
src/models/voxcpm/config.rs | 9 +
src/models/voxcpm/generate.rs | 47 +++--
src/models/voxcpm/minicpm4.rs | 58 ++++--
src/models/voxcpm/model.rs | 281 ++++++++++++++++++++---------
src/utils/mod.rs | 23 +++
tests/test_voxcpm2.rs | 49 +++++
tests/weight_test.rs | 23 +++
23 files changed, 548 insertions(+), 144 deletions(-)
create mode 100644 tests/test_voxcpm2.rs
diff --git a/Cargo.toml b/Cargo.toml
index 9556c7a..d9d1e23 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -4,7 +4,7 @@ version = "0.2.5"
edition = "2024"
repository = "https://github.com/jhqxxx/aha"
license = "Apache-2.0"
-description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM4, VoxCPM/1.5, DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)"
+description = "aha model inference library, now supports Qwen(2.5VL/3/3VL/3.5/ASR/3Embedding/3Reranker), MiniCPM4, VoxCPM(0.5B/1.5/2), DeepSeek-OCR/2, Hunyuan-OCR, PaddleOCR-VL/1.5, RMBG2.0, GLM(ASR-Nano-2512/OCR), Fun-ASR-Nano-2512, LFM(2/2.5/2VL/2.5VL)"
[dependencies]
candle-core = { version = "0.9.2" }
diff --git a/README.md b/README.md
index 8f3494e..ea589d1 100644
--- a/README.md
+++ b/README.md
@@ -33,7 +33,7 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
| **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5,
LFM2.5-VL, LFM2-VL |
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
-| **TTS** | VoxCPM, VoxCPM1.5 |
+| **TTS** | VoxCPM, VoxCPM1.5, VoxCPM2 |
| **Image** | RMBG-2.0 (background removal) |
| **Embedding** | Qwen3-Embedding, all-MiniLM-L6-v2 |
| **Reranker** | Qwen3-Reranker |
@@ -49,6 +49,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
## Changelog
+### 2026-04-08
+- add VoxCPM2
+
### 0.2.5 (2026-04-06)
- add qwen3-embedding/qwen3-reranker/all-minilm-l6-v2
diff --git a/README.zh-CN.md b/README.zh-CN.md
index a279ae5..854b3d4 100644
--- a/README.zh-CN.md
+++ b/README.zh-CN.md
@@ -32,7 +32,7 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
| **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5
LFM2.5-VL, LFM2-VL |
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2, PaddleOCR-VL,
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR |
| **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR |
-| **TTS** | VoxCPM, VoxCPM1.5 |
+| **TTS** | VoxCPM, VoxCPM1.5, VoxCPM2 |
| **图像** | RMBG-2.0 (背景移除) |
| **嵌入** | Qwen3-Embedding, all-MiniLM-L6-v2 |
| **重排序** | Qwen3-Reranker |
@@ -47,6 +47,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
- **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理
## 更新日志
+### 2026-04-08
+- 添加 VoxCPM2
+
## Changelog
### 0.2.5 (2026-04-06)
- 添加 qwen3-embedding/qwen3-reranker/all-minilm-l6-v2
diff --git a/docs/changelog.md b/docs/changelog.md
index 888d77f..21beaff 100644
--- a/docs/changelog.md
+++ b/docs/changelog.md
@@ -5,6 +5,9 @@ All notable changes to aha will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
+### 2026-04-08
+- add VoxCPM2
+
### 0.2.5 (2026-04-06)
- add qwen3-embedding/qwen3-reranker/all-minilm-l6-v2
diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md
index a4f28ce..64aaf5c 100644
--- a/docs/changelog.zh-CN.md
+++ b/docs/changelog.zh-CN.md
@@ -5,6 +5,9 @@
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
+### 2026-04-08
+- 添加 VoxCPM2
+
### 0.2.5 (2026-04-06)
- 添加 qwen3-embedding/qwen3-reranker/all-minilm-l6-v2
diff --git a/docs/model-card.zh-CN.md b/docs/model-card.zh-CN.md
index 2b78287..89f4640 100644
--- a/docs/model-card.zh-CN.md
+++ b/docs/model-card.zh-CN.md
@@ -74,3 +74,6 @@ extract structured information from documents. Prompts must follow a strict JSON
| 解析 | 1. Identify the formula in the image and represent it using LaTeX format.
2.Parse the table in the image into HTML.
3. Parse the chart in the image; use Mermaid format for flowcharts and Markdown for other charts.
4.Extract all information from the main body of the document image and represent it in markdown format, ignoring headers and footers. Tables should be expressed in HTML format, formulas in the document should be represented using LaTeX format, and the parsing should be organized according to the reading order. | 1. 识别图片中的公式,用 LaTeX 格式表示。
2. 把图中的表格解析为 HTML。
3. 解析图中的图表,对于流程图使用 Mermaid 格式表示,其他图表使用 Markdown 格式表示。
4. 提取文档图片中正文的所有信息用 markdown 格式表示,其中页眉、页脚部分忽略,表格用 html 格式表达,文档中公式用 latex 格式表示,按照阅读顺序组织进行解析。 |
| 信息提取 | 1. Output the value of Key.
2. Extract the content of the fields: ['key1','key2', ...] from the image and return it in JSON format.
3. Extract the subtitles from the image. | 1. 输出 Key 的值。
2. 提取图片中的: ['key1','key2', ...] 的字段内容,并按照 JSON 格式返回。
3. 提取图片中的字幕。 |
| 翻译 | First extract the text, then translate the text content into English. If it is a document, ignore the header and footer. Formulas should be represented in LaTeX format, and tables should be represented in HTML format. | 先提取文字,再将文字内容翻译为英文。若是文档,则其中页眉、页脚忽略。公式用latex格式表示,表格用html格式表示。 |
+
+# TTS
+## VoxCPM
\ No newline at end of file
diff --git a/docs/supported-models.md b/docs/supported-models.md
index 28b26c2..746b4ce 100644
--- a/docs/supported-models.md
+++ b/docs/supported-models.md
@@ -103,7 +103,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| Model | version | Model Id | License |
|-------|-----------|---------|---------|
-| **VoxCPM** | 1
1.5 | OpenBMB/VoxCPM-0.5B
OpenBMB/VoxCPM1.5 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
+| **VoxCPM** | 1
1.5
2 | OpenBMB/VoxCPM-0.5B
OpenBMB/VoxCPM1.5
OpenBMB/VoxCPM2 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Image Processing
diff --git a/docs/supported-models.zh-CN.md b/docs/supported-models.zh-CN.md
index 416f580..d48dda3 100644
--- a/docs/supported-models.zh-CN.md
+++ b/docs/supported-models.zh-CN.md
@@ -103,7 +103,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| 模型 | 版本 | 模型id | 开源协议 |
|------|--------|------|------|
-| **VoxCPM** | 1
1.5 | OpenBMB/VoxCPM-0.5B
OpenBMB/VoxCPM1.5 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
+| **VoxCPM** | 1
1.5
2 | OpenBMB/VoxCPM-0.5B
OpenBMB/VoxCPM1.5
OpenBMB/VoxCPM2 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## 图像处理
diff --git a/scripts/download_and_run.sh b/scripts/download_and_run.sh
index 81d92cb..30efafe 100755
--- a/scripts/download_and_run.sh
+++ b/scripts/download_and_run.sh
@@ -49,6 +49,7 @@ show_help() {
echo " AI-ModelScope/RMBG-2.0"
echo " OpenBMB/VoxCPM-0.5B"
echo " OpenBMB/VoxCPM1.5"
+ echo " OpenBMB/VoxCPM2"
echo " ZhipuAI/GLM-ASR-Nano-2512"
echo " FunAudioLLM/Fun-ASR-Nano-2512"
echo " ZhipuAI/GLM-OCR"
@@ -172,6 +173,9 @@ case $MODEL_ALIAS in
"OpenBMB/VoxCPM1.5")
MODEL_ID="openbmb/VoxCPM1.5"
;;
+ "OpenBMB/VoxCPM2")
+ MODEL_ID="openbmb/VoxCPM2"
+ ;;
"ZhipuAI/GLM-ASR-Nano-2512")
MODEL_ID="zai-org/GLM-ASR-Nano-2512"
;;
diff --git a/src/cli/mod.rs b/src/cli/mod.rs
index 7d963f1..27cf70a 100644
--- a/src/cli/mod.rs
+++ b/src/cli/mod.rs
@@ -301,7 +301,7 @@ pub(crate) fn run_run(args: RunArgs) -> anyhow::Result<()> {
WhichModel::RMBG2_0 => {
rmbg2_0::RMBG2_0Exec::run(&input, output.as_deref(), &weight_path)?;
}
- WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => {
+ WhichModel::VoxCPM | WhichModel::VoxCPM1_5 | WhichModel::VoxCPM2 => {
voxcpm::VoxCPMExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::GlmASRNano2512 => {
diff --git a/src/models/common/model_mapping.rs b/src/models/common/model_mapping.rs
index dd99fc2..7d38e37 100644
--- a/src/models/common/model_mapping.rs
+++ b/src/models/common/model_mapping.rs
@@ -74,6 +74,8 @@ pub enum WhichModel {
VoxCPM,
#[value(name = "OpenBMB/VoxCPM1.5")]
VoxCPM1_5,
+ #[value(name = "OpenBMB/VoxCPM2")]
+ VoxCPM2,
#[value(name = "ZhipuAI/GLM-ASR-Nano-2512")]
GlmASRNano2512,
#[value(name = "FunAudioLLM/Fun-ASR-Nano-2512")]
@@ -166,7 +168,7 @@ impl WhichModel {
// Image models
WhichModel::RMBG2_0 => "image",
// TTS models
- WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "tts",
+ WhichModel::VoxCPM | WhichModel::VoxCPM1_5 | WhichModel::VoxCPM2 => "tts",
WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B
| WhichModel::Qwen3Embedding8B
diff --git a/src/models/common/modules.rs b/src/models/common/modules.rs
index ab364f3..a384551 100644
--- a/src/models/common/modules.rs
+++ b/src/models/common/modules.rs
@@ -212,8 +212,8 @@ impl NaiveAttention {
pub fn forward_with_cache(
&mut self,
xs: &Tensor,
- cos: &Tensor,
- sin: &Tensor,
+ cos: Option<&Tensor>,
+ sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
tof32: bool,
) -> Result {
@@ -230,8 +230,15 @@ impl NaiveAttention {
let value_states = value_states
.reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
.transpose(1, 2)?;
- let (query_states, key_states) =
- apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?;
+ // let (query_states, key_states) =
+ // apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?;
+ let (query_states, key_states) = if let Some(cos) = cos
+ && let Some(sin) = sin
+ {
+ apply_rotary_pos_emb(&query_states, &key_states, cos, sin, tof32)?
+ } else {
+ (query_states, key_states)
+ };
let (key_states, value_states) = match &self.kv_cache {
None => (key_states, value_states),
Some((prev_k, prev_v)) => {
@@ -701,9 +708,9 @@ impl NaiveAttnGateUpDownMLPBlock {
) -> Result {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
- let xs = self
- .self_attn
- .forward_with_cache(&xs, cos, sin, attention_mask, false)?;
+ let xs =
+ self.self_attn
+ .forward_with_cache(&xs, Some(cos), Some(sin), attention_mask, false)?;
let residual = residual.add(&xs)?;
let xs = self.post_attention_layernorm.forward(&residual)?;
let xs = self.mlp.forward(&xs)?;
diff --git a/src/models/deepseek_ocr/model.rs b/src/models/deepseek_ocr/model.rs
index 5e9ded8..6539841 100644
--- a/src/models/deepseek_ocr/model.rs
+++ b/src/models/deepseek_ocr/model.rs
@@ -1018,9 +1018,9 @@ impl DeepseekV2DecoderLayer {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
- let xs = self
- .self_attn
- .forward_with_cache(&xs, cos, sin, attention_mask, false)?;
+ let xs =
+ self.self_attn
+ .forward_with_cache(&xs, Some(cos), Some(sin), attention_mask, false)?;
let residual = residual.add(&xs)?;
let xs = self.post_attention_layernorm.forward(&residual)?;
let xs = self.mlp.forward(&xs)?;
diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs
index a536ada..3f8d967 100644
--- a/src/models/minicpm4/model.rs
+++ b/src/models/minicpm4/model.rs
@@ -181,9 +181,9 @@ impl MiniCPMDecoderLayer {
) -> Result {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
- let xs = self
- .self_attn
- .forward_with_cache(&xs, cos, sin, attention_mask, true)?;
+ let xs =
+ self.self_attn
+ .forward_with_cache(&xs, Some(cos), Some(sin), attention_mask, true)?;
let xs = (residual
+ xs.affine(
self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(),
diff --git a/src/models/mod.rs b/src/models/mod.rs
index c21bf65..b2d0ec7 100644
--- a/src/models/mod.rs
+++ b/src/models/mod.rs
@@ -276,7 +276,7 @@ pub fn load_model<'a>(
let model = RMBG2_0Model::init(path, device, dtype)?;
ModelInstance::RMBG2_0(Box::new(model))
}
- WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => {
+ WhichModel::VoxCPM | WhichModel::VoxCPM1_5 | WhichModel::VoxCPM2 => {
let model = VoxCPMGenerate::init(path, device, dtype)?;
ModelInstance::VoxCPM(Box::new(model))
}
diff --git a/src/models/voxcpm/audio_vae.rs b/src/models/voxcpm/audio_vae.rs
index 1efafb0..0557e21 100644
--- a/src/models/voxcpm/audio_vae.rs
+++ b/src/models/voxcpm/audio_vae.rs
@@ -1,6 +1,11 @@
-use anyhow::{Ok, Result};
+use anyhow::{Result, anyhow};
use candle_core::{D, Tensor};
-use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder};
+use candle_nn::{
+ Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Embedding, Module, VarBuilder,
+ embedding,
+};
+
+use crate::utils::bucketize;
pub struct CausalConv1d {
conv1d: Conv1d,
@@ -398,12 +403,68 @@ impl CausalDecoderBlock {
}
}
+pub struct SampleRateConditionLayer {
+ cond_type: String,
+ scale_embed: Option,
+ bias_embed: Option,
+ cond_embed: Option,
+ // out_layer: Snake1d + WNCausalConv1d
+}
+
+impl SampleRateConditionLayer {
+ pub fn new(
+ vb: VarBuilder,
+ input_dim: usize,
+ sr_bin_buckets_len: usize,
+ cond_type: String,
+ // cond_dim: usize, // concat TODO
+ // out_layer: bool, //默认false
+ ) -> Result {
+ let (scale_embed, bias_embed, cond_embed) = if cond_type.contains("scale_bias") {
+ let scale_embed = embedding(sr_bin_buckets_len, input_dim, vb.pp("scale_embed"))?;
+ let bias_embed = embedding(sr_bin_buckets_len, input_dim, vb.pp("bias_embed"))?;
+ (Some(scale_embed), Some(bias_embed), None)
+ } else if cond_type.eq("add") {
+ let cond_embed = embedding(sr_bin_buckets_len, input_dim, vb.pp("cond_embed"))?;
+ (None, None, Some(cond_embed))
+ } else {
+ (None, None, None)
+ };
+ Ok(Self {
+ cond_type,
+ scale_embed,
+ bias_embed,
+ cond_embed,
+ })
+ }
+
+ pub fn forward(&self, x: &Tensor, sr_cond: &Tensor) -> Result {
+ if self.cond_type.contains("scale_bias")
+ && let Some(scale_embed) = &self.scale_embed
+ && let Some(bias_embed) = &self.bias_embed
+ {
+ Ok(
+ x.broadcast_mul(&scale_embed.forward(sr_cond)?.unsqueeze(D::Minus1)?)?
+ .broadcast_add(&bias_embed.forward(sr_cond)?.unsqueeze(D::Minus1)?)?,
+ )
+ } else if self.cond_type.eq("add")
+ && let Some(cond_embed) = &self.cond_embed
+ {
+ Ok(x.broadcast_add(&cond_embed.forward(sr_cond)?.unsqueeze(D::Minus1)?)?)
+ } else {
+ Err(anyhow!("not support cond_type"))
+ }
+ }
+}
+
pub struct CausalDecoder {
model0: WNCausalConv1d,
model1: WNCausalConv1d,
models: Vec,
model_minus_2: Snake1d,
model_minus_1: WNCausalConv1d,
+ sr_bin_boundaries: Option>,
+ sr_cond_model: Option>,
}
impl CausalDecoder {
@@ -414,6 +475,10 @@ impl CausalDecoder {
rates: Vec,
d_out: usize,
depthwise: bool,
+ sr_bin_boundaries: Option>,
+ cond_type: Option,
+ // cond_dim: Option,
+ // cond_out_layer: Option,
) -> Result {
let model0 = WNCausalConv1d::new(
vb.pp("model.0"),
@@ -429,8 +494,10 @@ impl CausalDecoder {
let vb_model = vb.pp("model");
let mut output_dim = channels;
let mut models = Vec::new();
+ let mut input_channels_vec = vec![];
for (i, stride) in rates.iter().enumerate() {
let input_dim = channels / 2_usize.pow(i as u32);
+ input_channels_vec.push(input_dim);
output_dim = channels / 2_usize.pow((i + 1) as u32);
let groups = if depthwise { output_dim } else { 1 };
let model_i = CausalDecoderBlock::new(
@@ -446,20 +513,53 @@ impl CausalDecoder {
let model_minus_2 = Snake1d::new(vb_model.pp(idx), output_dim)?;
let model_minus_1 =
WNCausalConv1d::new(vb_model.pp(idx + 1), output_dim, d_out, 7, 1, 3, 1, 1)?;
+ let (sr_cond_model, sr_bin_boundaries) = if let Some(sr) = sr_bin_boundaries
+ && let Some(cond_type) = cond_type
+ {
+ let sr_len = sr.len() + 1;
+ let vb_sr = vb.pp("sr_cond_model");
+ let mut sr_cond_model = vec![];
+ for (i, &input_dim) in input_channels_vec.iter().enumerate() {
+ let layer = SampleRateConditionLayer::new(
+ vb_sr.pp(i + 2),
+ input_dim,
+ sr_len,
+ cond_type.clone(),
+ )?;
+ sr_cond_model.push(layer);
+ }
+ (Some(sr_cond_model), Some(sr))
+ } else {
+ (None, None)
+ };
Ok(Self {
model0,
model1,
models,
model_minus_2,
model_minus_1,
+ sr_bin_boundaries,
+ sr_cond_model,
})
}
- pub fn forward(&self, x: &Tensor) -> Result {
+ pub fn forward(&self, x: &Tensor, sr_cond: Option) -> Result {
let x = self.model0.forward(x)?;
let mut x = self.model1.forward(&x)?;
- for model_i in &self.models {
- x = model_i.forward(&x)?;
+ if let Some(sr_cond) = sr_cond
+ && let Some(sr_models) = &self.sr_cond_model
+ && let Some(boundires) = &self.sr_bin_boundaries
+ {
+ let sr = bucketize(sr_cond, boundires)?;
+ let sr_cond = Tensor::new(vec![sr as u32], x.device())?;
+ for (model_i, sr_model_i) in self.models.iter().zip(sr_models.iter()) {
+ x = sr_model_i.forward(&x, &sr_cond)?;
+ x = model_i.forward(&x)?;
+ }
+ } else {
+ for model_i in &self.models {
+ x = model_i.forward(&x)?;
+ }
}
let x = self.model_minus_2.forward(&x)?;
let x = self.model_minus_1.forward(&x)?;
@@ -479,6 +579,8 @@ pub struct AudioVAE {
decoder: CausalDecoder,
pub sample_rate: usize,
pub chunk_size: usize,
+ sr_bin_boundaries: Option>,
+ out_sample_rate: usize,
}
impl AudioVAE {
@@ -490,6 +592,11 @@ impl AudioVAE {
decoder_dim: usize,
decoder_rates: Vec,
sample_rate: usize,
+ out_sample_rate: usize,
+ sr_bin_boundaries: Option>,
+ cond_type: Option,
+ // cond_dim: Option,
+ // cond_out_layer: Option,
) -> Result {
let latent_dim = match laten_dim {
Some(d) => d,
@@ -510,6 +617,10 @@ impl AudioVAE {
decoder_rates.clone(),
1,
true,
+ sr_bin_boundaries.clone(),
+ cond_type,
+ // cond_dim,
+ // cond_out_layer,
)?;
let chunk_size = hop_length;
Ok(Self {
@@ -522,7 +633,9 @@ impl AudioVAE {
encoder,
decoder,
sample_rate,
+ out_sample_rate,
chunk_size,
+ sr_bin_boundaries,
})
}
@@ -539,8 +652,13 @@ impl AudioVAE {
Ok(audio_data)
}
- pub fn decode(&self, z: &Tensor) -> Result {
- let x = self.decoder.forward(z)?;
+ pub fn decode(&self, z: &Tensor, sr_cond: Option) -> Result {
+ let sr_cond = if sr_cond.is_none() && self.sr_bin_boundaries.is_some() {
+ Some(self.out_sample_rate)
+ } else {
+ sr_cond
+ };
+ let x = self.decoder.forward(z, sr_cond)?;
Ok(x)
}
diff --git a/src/models/voxcpm/config.rs b/src/models/voxcpm/config.rs
index 17466f4..a4db69b 100644
--- a/src/models/voxcpm/config.rs
+++ b/src/models/voxcpm/config.rs
@@ -18,12 +18,14 @@ pub struct VoxMiniCPM4Config {
pub num_key_value_heads: usize,
pub rms_norm_eps: f64,
pub rope_theta: f32,
+ pub kv_channels: Option,
pub rope_scaling: VoxRopeScalingConfig,
pub vocab_size: usize,
pub scale_emb: f32,
pub dim_model_base: usize,
pub scale_depth: f32,
pub use_mup: bool,
+ pub no_rope: Option,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
@@ -32,6 +34,7 @@ pub struct VoxCPMEncoderConfig {
pub ffn_dim: usize,
pub num_heads: usize,
pub num_layers: usize,
+ pub kv_channels: Option,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
@@ -48,6 +51,8 @@ pub struct VoxCPMDitConfig {
pub ffn_dim: usize,
pub num_heads: usize,
pub num_layers: usize,
+ pub kv_channels: Option,
+ pub mean_mode: Option,
pub cfm_config: CfmConfig,
}
@@ -58,17 +63,21 @@ pub struct AudioVaeConfig {
pub latent_dim: usize,
pub decoder_dim: usize,
pub decoder_rates: Vec,
+ pub sr_bin_boundaries: Option>,
pub sample_rate: usize,
+ pub out_sample_rate: Option,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct VoxCPMConfig {
+ pub architecture: String,
pub lm_config: VoxMiniCPM4Config,
pub patch_size: usize,
pub feat_dim: usize,
pub scalar_quantization_latent_dim: usize,
pub scalar_quantization_scale: usize,
pub residual_lm_num_layers: usize,
+ pub residual_lm_no_rope: Option,
pub encoder_config: VoxCPMEncoderConfig,
pub dit_config: VoxCPMDitConfig,
pub audio_vae_config: Option,
diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs
index b12cb08..759fc1d 100644
--- a/src/models/voxcpm/generate.rs
+++ b/src/models/voxcpm/generate.rs
@@ -3,7 +3,7 @@ use std::collections::HashMap;
use crate::params::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
-use anyhow::{Ok, Result};
+use anyhow::{Result, anyhow};
use base64::{Engine, prelude::BASE64_STANDARD};
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
@@ -29,7 +29,7 @@ use crate::{
pub struct VoxCPMGenerate {
voxcpm: VoxCPMModel,
prompt_cache: Option>,
- sample_rate: usize,
+ out_sample_rate: usize,
model_name: String,
}
@@ -39,14 +39,12 @@ impl VoxCPMGenerate {
let config_path = path.to_string() + "/config.json";
let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
let model_list = find_type_files(path, "pth")?;
- // println!(" pth model_list: {:?}", model_list);
let mut dict_to_hashmap = HashMap::new();
let mut vae_dtype = candle_core::DType::F32;
for m in model_list {
let dict = read_all_with_key(m, Some("state_dict"))?;
vae_dtype = dict[0].1.dtype();
for (k, v) in dict {
- // println!("key: {}, tensor shape: {:?}", k, v);
dict_to_hashmap.insert(k, v);
}
}
@@ -60,6 +58,8 @@ impl VoxCPMGenerate {
decoder_dim: 1536,
decoder_rates: vec![8, 8, 5, 2],
sample_rate: 16000,
+ out_sample_rate: None,
+ sr_bin_boundaries: None,
},
};
let model_name = std::path::Path::new(path)
@@ -67,11 +67,6 @@ impl VoxCPMGenerate {
.and_then(|s| s.to_str())
.unwrap_or("VoxCPM")
.to_string();
- // let model_name = if audio_config.sample_rate == 16000 {
- // "VoxCPM".to_string()
- // } else {
- // "VoxCPM1.5".to_string()
- // };
let audio_vae = AudioVAE::new(
vb_vae,
audio_config.encoder_dim,
@@ -80,6 +75,13 @@ impl VoxCPMGenerate {
audio_config.decoder_dim,
audio_config.decoder_rates.clone(),
audio_config.sample_rate,
+ audio_config
+ .out_sample_rate
+ .unwrap_or(audio_config.sample_rate),
+ audio_config.sr_bin_boundaries,
+ Some("scale_bias".to_string()),
+ // Some(128),
+ // Some(false),
)?;
let cfg_dtype = config.dtype.as_str();
@@ -105,11 +107,13 @@ impl VoxCPMGenerate {
};
let tokenizer = SingleChineseTokenizer::new(path)?;
let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?;
-
+ let out_sample_rate = audio_config
+ .out_sample_rate
+ .unwrap_or(audio_config.sample_rate);
Ok(Self {
voxcpm,
prompt_cache: None,
- sample_rate: audio_config.sample_rate,
+ out_sample_rate,
model_name,
})
}
@@ -208,13 +212,15 @@ impl VoxCPMGenerate {
}
pub fn sample_rate(&self) -> usize {
- self.sample_rate
+ self.out_sample_rate
}
}
impl GenerateModel for VoxCPMGenerate {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
let prompt_text = extract_metadata_value::(&mes.metadata, "prompt_text");
+ let control_instruction =
+ extract_metadata_value::(&mes.metadata, "control_instruction");
let min_len = extract_metadata_value::(&mes.metadata, "min_len").unwrap_or(2);
let max_len = extract_metadata_value::(&mes.metadata, "max_len").unwrap_or(4096);
let inference_timesteps =
@@ -223,13 +229,26 @@ impl GenerateModel for VoxCPMGenerate {
let retry_badcase_ratio_threshold =
extract_metadata_value::(&mes.metadata, "retry_badcase_ratio_threshold")
.unwrap_or(6.0);
- let target_text = extract_user_text(&mes)?;
+
let prompt_wav = extract_audio_url(&mes);
let prompt_wav_path = if !prompt_wav.is_empty() {
Some(prompt_wav[0].clone())
} else {
None
};
+ if !self.model_name.contains("2") && prompt_wav_path.is_some() && prompt_text.is_none() {
+ return Err(anyhow!(
+ "reference mode is only supported with VoxCPM2 models"
+ ));
+ }
+ let mut target_text = extract_user_text(&mes)?;
+ if let Some(instruction) = control_instruction
+ && self.model_name.contains("2")
+ && prompt_text.is_none()
+ && prompt_wav_path.is_none()
+ {
+ target_text = format!("({instruction}){target_text}");
+ }
let audio = self
.voxcpm
.generate(
@@ -245,7 +264,7 @@ impl GenerateModel for VoxCPMGenerate {
.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.out_sample_rate as u32)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name);
self.voxcpm.clear_kv_cache();
diff --git a/src/models/voxcpm/minicpm4.rs b/src/models/voxcpm/minicpm4.rs
index cd38f19..e656a53 100644
--- a/src/models/voxcpm/minicpm4.rs
+++ b/src/models/voxcpm/minicpm4.rs
@@ -25,7 +25,9 @@ pub struct MiniCPMLongRoPE {
}
impl MiniCPMLongRoPE {
pub fn new(cfg: &VoxMiniCPM4Config, device: &Device, dtype: DType) -> Result {
- let head_dim = cfg.hidden_size / cfg.num_attention_heads;
+ let head_dim = cfg
+ .kv_channels
+ .unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
let rope_theta = cfg.rope_theta;
let short_factor = cfg.rope_scaling.short_factor.clone();
let long_factor = cfg.rope_scaling.short_factor.clone();
@@ -117,7 +119,10 @@ impl MiniCPMDecoderLayer {
cfg.hidden_size,
cfg.num_attention_heads,
cfg.num_key_value_heads,
- None,
+ Some(
+ cfg.kv_channels
+ .unwrap_or(cfg.hidden_size / cfg.num_attention_heads),
+ ),
false,
None,
None,
@@ -155,15 +160,15 @@ impl MiniCPMDecoderLayer {
pub fn forward(
&self,
xs: &Tensor,
- cos: &Tensor,
- sin: &Tensor,
+ cos: Option<&Tensor>,
+ sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
) -> Result {
let residual = xs.clone();
let xs = self.input_layernorm.forward(xs)?;
let xs = self
.self_attn
- .forward(&xs, Some(cos), Some(sin), attention_mask, true)?;
+ .forward(&xs, cos, sin, attention_mask, true)?;
let xs = if self.use_mup {
(residual
+ xs.affine(
@@ -191,8 +196,8 @@ impl MiniCPMDecoderLayer {
pub fn forward_with_cache(
&mut self,
xs: &Tensor,
- cos: &Tensor,
- sin: &Tensor,
+ cos: Option<&Tensor>,
+ sin: Option<&Tensor>,
attention_mask: Option<&Tensor>,
) -> Result {
let residual = xs.clone();
@@ -232,7 +237,7 @@ pub struct MiniCPMModel {
pub embed_tokens: Option,
layers: Vec,
norm: RmsNorm,
- rope_emb: MiniCPMLongRoPE,
+ rope_emb: Option,
}
impl MiniCPMModel {
@@ -255,7 +260,15 @@ impl MiniCPMModel {
layers.push(layer);
}
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
- let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?;
+
+ let rope_emb = if let Some(flag) = cfg.no_rope
+ && flag
+ {
+ None
+ } else {
+ Some(MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?)
+ };
+
Ok(Self {
// cfg,
embed_tokens,
@@ -284,11 +297,20 @@ impl MiniCPMModel {
)?)
}
};
- let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
+ let (cos, sin) = if let Some(rope_emb) = &mut self.rope_emb {
+ let (cos, sin) = rope_emb.forward(position_id, seq_len)?;
+ (Some(cos), Some(sin))
+ } else {
+ (None, None)
+ };
let mut hidden_states = input_embeds.clone();
for decode_layer in &self.layers {
- hidden_states =
- decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?;
+ hidden_states = decode_layer.forward(
+ &hidden_states,
+ cos.as_ref(),
+ sin.as_ref(),
+ attention_mask.as_ref(),
+ )?;
}
hidden_states = self.norm.forward(&hidden_states)?;
Ok(hidden_states)
@@ -317,13 +339,19 @@ impl MiniCPMModel {
)?)
}
};
- let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
+ // let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?;
+ let (cos, sin) = if let Some(rope_emb) = &mut self.rope_emb {
+ let (cos, sin) = rope_emb.forward(position_id, seq_len)?;
+ (Some(cos), Some(sin))
+ } else {
+ (None, None)
+ };
let mut hidden_states = input_embeds.clone();
for decode_layer in &mut self.layers {
hidden_states = decode_layer.forward_with_cache(
&hidden_states,
- &cos,
- &sin,
+ cos.as_ref(),
+ sin.as_ref(),
attention_mask.as_ref(),
)?;
}
diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs
index d90ff7a..16265b0 100644
--- a/src/models/voxcpm/model.rs
+++ b/src/models/voxcpm/model.rs
@@ -70,7 +70,6 @@ impl SinusoidalPosEmb {
.affine(-dif, 0.0)?
.exp()?
.to_dtype(x.dtype())?;
-
let emb = x
.unsqueeze(1)?
.contiguous()?
@@ -118,6 +117,7 @@ pub struct VoxCPMLocDiT {
time_mlp: TimestepEmbedding,
delta_time_mlp: TimestepEmbedding,
decoder: MiniCPMModel,
+ version: usize,
// config: VoxMiniCPM4Config,
// in_channels: usize,
}
@@ -142,6 +142,11 @@ impl VoxCPMLocDiT {
)?;
assert_eq!(config.vocab_size, 0, "vocab_size must be 0 for local DiT");
let decoder = MiniCPMModel::new(vb.pp("decoder"), config.clone())?;
+ let version = if config.kv_channels.is_some() {
+ 2usize
+ } else {
+ 1
+ };
Ok(Self {
in_proj,
cond_proj,
@@ -150,6 +155,7 @@ impl VoxCPMLocDiT {
time_mlp,
delta_time_mlp,
decoder,
+ version,
// config,
// in_channels,
})
@@ -176,11 +182,19 @@ impl VoxCPMLocDiT {
.to_dtype(x.dtype())?;
let dt = self.delta_time_mlp.forward(&dt)?;
let t = t.add(&dt)?;
-
- let x = Tensor::cat(&[mu.add(&t)?.unsqueeze(1)?, cond, x], 1)?;
- let hidden = self.decoder.forward(&x, 0, false)?;
- let select_len = hidden.dims()[1] - (prefix + 1);
- let hidden = hidden.narrow(1, prefix + 1, select_len)?;
+ let hidden = if self.version == 2 {
+ let (b, _, dim) = x.dims3()?;
+ let mu = mu.reshape((b, (), dim))?;
+ let x = Tensor::cat(&[&mu, &t.unsqueeze(1)?, &cond, &x], 1)?;
+ let hidden = self.decoder.forward(&x, 0, false)?;
+ let select_len = hidden.dim(1)? - (prefix + mu.dim(1)? + 1);
+ hidden.narrow(1, prefix + mu.dim(1)? + 1, select_len)?
+ } else {
+ let x = Tensor::cat(&[mu.add(&t)?.unsqueeze(1)?, cond, x], 1)?;
+ let hidden = self.decoder.forward(&x, 0, false)?;
+ let select_len = hidden.dim(1)? - (prefix + 1);
+ hidden.narrow(1, prefix + 1, select_len)?
+ };
let hidden = self.out_proj.forward(&hidden)?;
let hidden = hidden.transpose(1, 2)?.contiguous()?;
Ok(hidden)
@@ -194,6 +208,7 @@ pub struct UnifiedCFM {
in_channels: usize,
mean_mode: bool,
estimator: VoxCPMLocDiT,
+ // architecture: String,
}
impl UnifiedCFM {
@@ -202,6 +217,7 @@ impl UnifiedCFM {
_cfm_params: CfmConfig,
estimator: VoxCPMLocDiT,
mean_mode: bool,
+ // architecture: String,
) -> Result {
// let solver = cfm_params.solver;
// let sigma_min = cfm_params.sigma_min;
@@ -213,6 +229,7 @@ impl UnifiedCFM {
in_channels,
mean_mode,
estimator,
+ // architecture,
})
}
@@ -233,6 +250,7 @@ impl UnifiedCFM {
let z = Tensor::randn(0.0f32, 1.0, (b, self.in_channels, t), mu.device())?
.to_dtype(dtype)?
.affine(temperature, 0.0)?;
+ // let z = Tensor::ones((b, self.in_channels, t), dtype, mu.device())?;
let t_span = linspace(1.0, 0.0, n_timesteps + 1, mu.device())?.to_dtype(dtype)?;
let t_span = t_span
.affine(f64::consts::PI / 2.0, 0.0)?
@@ -362,8 +380,10 @@ impl VoxCPMLocEnc {
pub struct VoxCPMModel {
config: VoxCPMConfig,
patch_size: usize,
- audio_start_token: usize,
- // audio_end_token: usize,
+ audio_start_token: u32,
+ // audio_end_token: u32,
+ ref_audio_start_token: u32,
+ ref_audio_end_token: u32,
chunk_size: usize,
sample_rate: usize,
tokenizer: SingleChineseTokenizer,
@@ -376,6 +396,7 @@ pub struct VoxCPMModel {
enc_to_lm_proj: Linear,
lm_to_dit_proj: Linear,
res_to_dit_proj: Linear,
+ fusion_concat_proj: Option,
stop_proj: Linear,
stop_head: Linear,
device: Device,
@@ -390,17 +411,17 @@ impl VoxCPMModel {
audio_vae: AudioVAE,
) -> Result {
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
- let audio_start_token = 101usize;
- // let audio_end_token = 102usize;
let mut residual_lm_config = config.lm_config.clone();
residual_lm_config.num_hidden_layers = config.residual_lm_num_layers;
residual_lm_config.vocab_size = 0;
+ residual_lm_config.no_rope = config.residual_lm_no_rope;
let residual_lm = MiniCPMModel::new(vb.pp("residual_lm"), residual_lm_config)?;
let mut encoder_config = config.lm_config.clone();
encoder_config.hidden_size = config.encoder_config.hidden_dim;
encoder_config.intermediate_size = config.encoder_config.ffn_dim;
encoder_config.num_attention_heads = config.encoder_config.num_heads;
encoder_config.num_hidden_layers = config.encoder_config.num_layers;
+ encoder_config.kv_channels = config.encoder_config.kv_channels;
encoder_config.vocab_size = 0;
let feat_encoder =
VoxCPMLocEnc::new(vb.pp("feat_encoder"), encoder_config, config.feat_dim)?;
@@ -410,6 +431,7 @@ impl VoxCPMModel {
decoder_config.intermediate_size = config.dit_config.ffn_dim;
decoder_config.num_attention_heads = config.dit_config.num_heads;
decoder_config.num_hidden_layers = config.dit_config.num_layers;
+ decoder_config.kv_channels = config.dit_config.kv_channels;
decoder_config.vocab_size = 0;
let estimator = VoxCPMLocDiT::new(
vb.pp("feat_decoder.estimator"),
@@ -421,6 +443,7 @@ impl VoxCPMModel {
config.dit_config.cfm_config.clone(),
estimator,
false,
+ // config.architecture.clone(),
)?;
let fsq_layer = ScalarQuantizationLayer::new(
vb.pp("fsq_layer"),
@@ -445,6 +468,16 @@ impl VoxCPMModel {
vb.pp("res_to_dit_proj"),
)?;
+ let fusion_concat_proj = if config.architecture.to_lowercase().eq("voxcpm2") {
+ Some(linear(
+ config.lm_config.hidden_size * 2,
+ config.lm_config.hidden_size,
+ vb.pp("fusion_concat_proj"),
+ )?)
+ } else {
+ None
+ };
+
let stop_proj = linear(
config.lm_config.hidden_size,
config.lm_config.hidden_size,
@@ -456,8 +489,10 @@ impl VoxCPMModel {
Ok(Self {
config,
patch_size,
- audio_start_token,
- // audio_end_token,
+ audio_start_token: 101,
+ // audio_end_token: 102,
+ ref_audio_start_token: 103,
+ ref_audio_end_token: 104,
chunk_size: audio_vae.chunk_size,
sample_rate: audio_vae.sample_rate,
tokenizer,
@@ -470,6 +505,7 @@ impl VoxCPMModel {
enc_to_lm_proj,
lm_to_dit_proj,
res_to_dit_proj,
+ fusion_concat_proj,
stop_proj,
stop_head,
device: vb.device().clone(),
@@ -489,72 +525,128 @@ impl VoxCPMModel {
// retry_badcase: bool,
retry_badcase_ratio_threshold: f64,
) -> Result {
- let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path {
- None => {
- let text_token = self.tokenizer.encode(target_text.clone())?;
- let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
- let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?;
- let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
- let text_length = text_token.dim(0)?;
- let audio_feat = Tensor::zeros(
- (text_length, self.patch_size, self.audio_vae.latent_dim),
- DType::F32,
- &self.device,
- )?;
- let text_mask = Tensor::ones(text_length, self.dtype, &self.device)?;
- let audio_mask = Tensor::zeros(text_length, self.dtype, &self.device)?;
- (text_token, text_mask, audio_feat, audio_mask)
+ let (text_token, text_mask, audio_feat, audio_mask) = if let Some(prompt_text) = prompt_text
+ && let Some(path) = prompt_wav_path
+ {
+ let text = prompt_text + &target_text;
+ let text_token = self.tokenizer.encode(text)?;
+ let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
+ let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
+ let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
+ let text_length = text_token.dim(0)?;
+ let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
+ let patch_len = self.patch_size * self.chunk_size;
+ if audio.dim(1)? % patch_len != 0 {
+ audio =
+ audio.pad_with_zeros(D::Minus1, patch_len - audio.dim(1)? % patch_len, 0)?;
}
- Some(path) => {
- let text = prompt_text.unwrap_or("".to_string()) + &target_text;
- let text_token = self.tokenizer.encode(text)?;
- let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
- let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?;
- let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
- let text_length = text_token.dim(0)?;
- let mut audio =
- load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
- let patch_len = self.patch_size * self.chunk_size;
- if audio.dim(1)? % patch_len != 0 {
- audio = audio.pad_with_zeros(
- D::Minus1,
- // 0,
- // patch_len - audio.dim(1)? % patch_len,
- patch_len - audio.dim(1)? % patch_len,
- 0,
- )?;
- }
- let audio_feat = self.audio_vae.encode(&audio, Some(self.sample_rate))?;
- let audio_feat = audio_feat
- .reshape((self.audio_vae.latent_dim, (), self.patch_size))?
- .permute((1, 2, 0))?;
- // let dim0 = audio_feat.dim(0)? - 1;
- // let audio_feat = audio_feat.i(..dim0)?;
- let audio_length = audio_feat.dim(0)?;
- let text_pad_token = Tensor::zeros(audio_length, DType::U32, &self.device)?;
- let text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
- let audio_pad_feat = Tensor::zeros(
- (text_length, self.patch_size, self.audio_vae.latent_dim),
- audio_feat.dtype(),
- &self.device,
- )?;
- let audio_feat = Tensor::cat(&[audio_pad_feat, audio_feat], 0)?;
- let text_mask = Tensor::cat(
- &[
- Tensor::ones(text_length, self.dtype, &self.device)?,
- Tensor::zeros(audio_length, self.dtype, &self.device)?,
- ],
- D::Minus1,
- )?;
- let audio_mask = Tensor::cat(
- &[
- Tensor::zeros(text_length, self.dtype, &self.device)?,
- Tensor::ones(audio_length, self.dtype, &self.device)?,
- ],
- D::Minus1,
- )?;
- (text_token, text_mask, audio_feat, audio_mask)
+ let audio_feat = self.audio_vae.encode(&audio, Some(self.sample_rate))?;
+ let audio_feat = audio_feat
+ .reshape((self.audio_vae.latent_dim, (), self.patch_size))?
+ .permute((1, 2, 0))?;
+ let audio_length = audio_feat.dim(0)?;
+ let text_pad_token = Tensor::zeros(audio_length, DType::U32, &self.device)?;
+ let text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
+ let audio_pad_feat = Tensor::zeros(
+ (text_length, self.patch_size, self.audio_vae.latent_dim),
+ audio_feat.dtype(),
+ &self.device,
+ )?;
+ let audio_feat = Tensor::cat(&[audio_pad_feat, audio_feat], 0)?;
+ let text_mask = Tensor::cat(
+ &[
+ Tensor::ones(text_length, self.dtype, &self.device)?,
+ Tensor::zeros(audio_length, self.dtype, &self.device)?,
+ ],
+ D::Minus1,
+ )?;
+ let audio_mask = Tensor::cat(
+ &[
+ Tensor::zeros(text_length, self.dtype, &self.device)?,
+ Tensor::ones(audio_length, self.dtype, &self.device)?,
+ ],
+ D::Minus1,
+ )?;
+ (text_token, text_mask, audio_feat, audio_mask)
+ } else if let Some(path) = prompt_wav_path {
+ let text_token = self.tokenizer.encode(target_text.clone())?;
+ let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
+ let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
+ let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
+ let text_length = text_token.dim(0)?;
+ let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
+ let patch_len = self.patch_size * self.chunk_size;
+ if audio.dim(1)? % patch_len != 0 {
+ audio =
+ audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
}
+ let audio_feat = self.audio_vae.encode(&audio, Some(self.sample_rate))?;
+ let audio_feat = audio_feat
+ .reshape((self.audio_vae.latent_dim, (), self.patch_size))?
+ .permute((1, 2, 0))?;
+ let ref_len = audio_feat.dim(0)?;
+ let z1 = Tensor::zeros(
+ (1, self.patch_size, self.audio_vae.latent_dim),
+ DType::F32,
+ &self.device,
+ )?;
+ let ref_start = Tensor::new(vec![self.ref_audio_start_token], &self.device)?;
+ let ref_end = Tensor::new(vec![self.ref_audio_end_token], &self.device)?;
+ let ref_token = Tensor::zeros(ref_len, DType::U32, &self.device)?;
+ let ref_tokens = Tensor::cat(&[&ref_start, &ref_token, &ref_end], 0)?;
+ let feats = Tensor::cat(&[&z1, &audio_feat, &z1], 0)?;
+ let t_mask = Tensor::cat(
+ &[
+ Tensor::new(vec![1.0f32], &self.device)?.to_dtype(self.dtype)?,
+ Tensor::zeros(ref_len, self.dtype, &self.device)?,
+ Tensor::new(vec![1.0f32], &self.device)?.to_dtype(self.dtype)?,
+ ],
+ 0,
+ )?;
+ let a_mask = Tensor::cat(
+ &[
+ Tensor::new(vec![0.0f32], &self.device)?.to_dtype(self.dtype)?,
+ Tensor::ones(ref_len, self.dtype, &self.device)?,
+ Tensor::new(vec![0.0f32], &self.device)?.to_dtype(self.dtype)?,
+ ],
+ 0,
+ )?;
+ let text_pad_feat = Tensor::zeros(
+ (text_length, self.patch_size, self.audio_vae.latent_dim),
+ DType::F32,
+ &self.device,
+ )?;
+ let text_token = Tensor::cat(&[&ref_tokens, &text_token], 0)?;
+ let audio_feat = Tensor::cat(&[&feats, &text_pad_feat], 0)?;
+ let text_mask = Tensor::cat(
+ &[
+ &t_mask,
+ &Tensor::ones(text_length, self.dtype, &self.device)?,
+ ],
+ 0,
+ )?;
+ let audio_mask = Tensor::cat(
+ &[
+ &a_mask,
+ &Tensor::zeros(text_length, self.dtype, &self.device)?,
+ ],
+ 0,
+ )?;
+ (text_token, text_mask, audio_feat, audio_mask)
+ } else {
+ let text_token = self.tokenizer.encode(target_text.clone())?;
+ let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
+ let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
+ let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
+ let text_length = text_token.dim(0)?;
+ let audio_feat = Tensor::zeros(
+ (text_length, self.patch_size, self.audio_vae.latent_dim),
+ DType::F32,
+ &self.device,
+ )?;
+ let text_mask = Tensor::ones(text_length, self.dtype, &self.device)?;
+ let audio_mask = Tensor::zeros(text_length, self.dtype, &self.device)?;
+ (text_token, text_mask, audio_feat, audio_mask)
};
let target_text_length = self.tokenizer.encode(target_text)?.len();
// let max_len = if retry_badcase {
@@ -605,7 +697,7 @@ impl VoxCPMModel {
)?;
let decode_audio = self
.audio_vae
- .decode(&latent_pred.to_dtype(DType::F32)?)?
+ .decode(&latent_pred.to_dtype(DType::F32)?, None)?
.squeeze(1)?;
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640;
let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?;
@@ -645,6 +737,9 @@ impl VoxCPMModel {
.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?;
let mut prefix_feat_cond = feat.i((.., t - 1, ..))?;
let mut pred_feat_seq = Vec::new();
+ // if feat_mask.i((1, t-1))?.to_scalar::()? == 0.0 {
+ // // TODO for stream
+ // }
let mut position_id = 0;
let mut seq_len = t;
let enc_outputs = self
@@ -655,20 +750,27 @@ impl VoxCPMModel {
.forward(&enc_outputs)?
.broadcast_mul(&feat_mask.unsqueeze(D::Minus1)?)?
.add(&enc_outputs.broadcast_mul(&text_mask.unsqueeze(D::Minus1)?)?)?;
-
let mut lm_hidden = enc_outputs.i((.., t - 1, ..))?;
-
- let input_embeds =
- enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?;
+ let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
+ let feat = feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?;
+ let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
+ fusion.forward(&concat)?
+ } else {
+ enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?
+ };
let residual_enc_outputs = self
.residual_lm
.forward_with_cache(&input_embeds, position_id)?;
let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?;
-
for i in 0..max_len {
let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit]
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit]
- let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
+ // let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
+ let dit_hidden = if self.fusion_concat_proj.is_some() {
+ Tensor::cat(&[&dit_hidden_1, &dit_hidden_2], D::Minus1)?
+ } else {
+ dit_hidden_1.add(&dit_hidden_2)?
+ };
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
let pred_feat = self
.feat_decoder
@@ -705,9 +807,16 @@ impl VoxCPMModel {
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
.squeeze(1)?;
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
+ let curr_residual_input = if let Some(fusion) = &self.fusion_concat_proj {
+ let curr_embed = curr_embed.i((.., 0, ..))?;
+ let concat = Tensor::cat(&[&lm_hidden, &curr_embed], D::Minus1)?;
+ fusion.forward(&concat)?
+ } else {
+ lm_hidden.add(&curr_embed.i((.., 0, ..))?)?
+ };
residual_hidden = self
.residual_lm
- .forward_with_cache(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)?
+ .forward_with_cache(&curr_residual_input, position_id)?
.squeeze(1)?;
}
let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d)
@@ -716,8 +825,6 @@ impl VoxCPMModel {
.permute((0, 3, 1, 2))?
.reshape((b, d, ()))?
.contiguous()?;
- // self.base_lm.clear_kv_cache();
- // self.residual_lm.clear_kv_cache();
self.clear_kv_cache();
Ok(feat_pred)
}
@@ -770,7 +877,7 @@ impl VoxCPMModel {
Some(token) => Tensor::cat(&[token, &target_text_token], 0)?,
None => target_text_token,
};
- let audio_start = Tensor::new(vec![self.audio_start_token as u32], &self.device)?;
+ let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let text_length = text_token.dim(0)?;
let (audio_length, audio_feat) = match prompt_cache.get("audio_feat") {
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index ea1dc2f..d5325bf 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -698,6 +698,29 @@ pub fn bytes_to_human(bytes: u64) -> String {
}
}
+pub fn bucketize(input: usize, boundaries: &[usize]) -> Result {
+ if boundaries.is_empty() {
+ return Err(anyhow!("bucketize param boundaries can not be empty"));
+ }
+ match boundaries.binary_search(&input) {
+ Ok(i) => Ok(i),
+ Err(i) => Ok(i),
+ }
+ // let mut index = 0;
+ // let mut change = false;
+ // for i in 0..boundaries.len() {
+ // if input <= boundaries[i] {
+ // index = i;
+ // change = true;
+ // break;
+ // }
+ // }
+ // if !change {
+ // index = boundaries.len();
+ // }
+ // Ok(index)
+}
+
#[cfg(test)]
mod tests {
use super::*;
diff --git a/tests/test_voxcpm2.rs b/tests/test_voxcpm2.rs
new file mode 100644
index 0000000..2981db9
--- /dev/null
+++ b/tests/test_voxcpm2.rs
@@ -0,0 +1,49 @@
+use std::time::Instant;
+
+use aha::{
+ models::{GenerateModel, voxcpm::generate::VoxCPMGenerate},
+ params::chat::ChatCompletionParameters,
+ utils::audio_utils::extract_and_save_audio_from_response,
+};
+use anyhow::Result;
+
+#[test]
+fn voxcpm2_use_message_generate() -> Result<()> {
+ // RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm2 voxcpm2_use_message_generate -r -- --nocapture
+ // control_instruction
+ let save_dir =
+ aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
+ let model_path = format!("{}/OpenBMB/VoxCPM2/", save_dir);
+ let message = r#"
+ {
+ "model": "OpenBMB/VoxCPM2",
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "老板儿,来碗担担面,多放海椒,再加个煎蛋哈"
+ }
+ ]
+ }
+ ],
+ "metadata": {"control_instruction": "萝莉音"}
+ }
+ "#;
+ let mes: ChatCompletionParameters = serde_json::from_str(message)?;
+ let i_start = Instant::now();
+ let mut voxcpm_generate = VoxCPMGenerate::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 generate = voxcpm_generate.generate(mes)?;
+ let save_path = extract_and_save_audio_from_response(&generate, "./")?;
+ for path in save_path {
+ println!("save audio: {}", path);
+ }
+ let i_duration = i_start.elapsed();
+ println!("Time elapsed in generate is: {:?}", i_duration);
+ Ok(())
+}
diff --git a/tests/weight_test.rs b/tests/weight_test.rs
index ebe4d4c..b81ff80 100644
--- a/tests/weight_test.rs
+++ b/tests/weight_test.rs
@@ -328,3 +328,26 @@ fn gguf_weight() -> Result<()> {
// }
Ok(())
}
+
+#[test]
+fn voxcpm2_weight() -> Result<()> {
+ // cargo test -F cuda --test weight_test voxcpm2_weight -r -- --nocapture
+ let save_dir =
+ aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
+ let model_path = format!("{}/OpenBMB/VoxCPM2/", save_dir);
+ let model_list = find_type_files(&model_path, "pth")?;
+ 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 {
+ println!("key: {}, tensor shape: {:?}", k, v);
+ // dict_to_hashmap.insert(k, v);
+ }
+ }
+
+ Ok(())
+}