add VoxCPM2

This commit is contained in:
jhqxxx
2026-04-08 18:55:03 +08:00
parent 00d6f34e93
commit adf31dc16f
23 changed files with 548 additions and 144 deletions
+1 -1
View File
@@ -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" }
+4 -1
View File
@@ -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, <br> LFM2.5-VL, LFM2-VL |
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL <br> 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
+4 -1
View File
@@ -32,7 +32,7 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
| **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5 <br> LFM2.5-VL, LFM2-VL |
| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2, PaddleOCR-VL, <br>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
+3
View File
@@ -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
+3
View File
@@ -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
+3
View File
@@ -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. <br> 2.Parse the table in the image into HTML. <br> 3. Parse the chart in the image; use Mermaid format for flowcharts and Markdown for other charts. <br> 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 格式表示。 <br> 2. 把图中的表格解析为 HTML。 <br> 3. 解析图中的图表,对于流程图使用 Mermaid 格式表示,其他图表使用 Markdown 格式表示。 <br> 4. 提取文档图片中正文的所有信息用 markdown 格式表示,其中页眉、页脚部分忽略,表格用 html 格式表达,文档中公式用 latex 格式表示,按照阅读顺序组织进行解析。 |
| 信息提取 | 1. Output the value of Key. <br> 2. Extract the content of the fields: ['key1','key2', ...] from the image and return it in JSON format. <br> 3. Extract the subtitles from the image. | 1. 输出 Key 的值。 <br> 2. 提取图片中的: ['key1','key2', ...] 的字段内容,并按照 JSON 格式返回。 <br> 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
+1 -1
View File
@@ -103,7 +103,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| Model | version | Model Id | License |
|-------|-----------|---------|---------|
| **VoxCPM** | 1<br>1.5 | OpenBMB/VoxCPM-0.5B <br> OpenBMB/VoxCPM1.5 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **VoxCPM** | 1<br>1.5<br>2 | OpenBMB/VoxCPM-0.5B <br> OpenBMB/VoxCPM1.5 <br> OpenBMB/VoxCPM2 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Image Processing
+1 -1
View File
@@ -103,7 +103,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| 模型 | 版本 | 模型id | 开源协议 |
|------|--------|------|------|
| **VoxCPM** | 1<br>1.5 | OpenBMB/VoxCPM-0.5B <br> OpenBMB/VoxCPM1.5 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **VoxCPM** | 1<br>1.5<br>2 | OpenBMB/VoxCPM-0.5B <br> OpenBMB/VoxCPM1.5 <br> OpenBMB/VoxCPM2 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## 图像处理
+4
View File
@@ -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"
;;
+1 -1
View File
@@ -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 => {
+3 -1
View File
@@ -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
+14 -7
View File
@@ -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<Tensor> {
@@ -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<Tensor> {
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)?;
+3 -3
View File
@@ -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)?;
+3 -3
View File
@@ -181,9 +181,9 @@ impl MiniCPMDecoderLayer {
) -> Result<Tensor> {
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(),
+1 -1
View File
@@ -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))
}
+125 -7
View File
@@ -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<Embedding>,
bias_embed: Option<Embedding>,
cond_embed: Option<Embedding>,
// 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<Self> {
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<Tensor> {
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<CausalDecoderBlock>,
model_minus_2: Snake1d,
model_minus_1: WNCausalConv1d,
sr_bin_boundaries: Option<Vec<usize>>,
sr_cond_model: Option<Vec<SampleRateConditionLayer>>,
}
impl CausalDecoder {
@@ -414,6 +475,10 @@ impl CausalDecoder {
rates: Vec<usize>,
d_out: usize,
depthwise: bool,
sr_bin_boundaries: Option<Vec<usize>>,
cond_type: Option<String>,
// cond_dim: Option<usize>,
// cond_out_layer: Option<bool>,
) -> Result<Self> {
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<Tensor> {
pub fn forward(&self, x: &Tensor, sr_cond: Option<usize>) -> Result<Tensor> {
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<Vec<usize>>,
out_sample_rate: usize,
}
impl AudioVAE {
@@ -490,6 +592,11 @@ impl AudioVAE {
decoder_dim: usize,
decoder_rates: Vec<usize>,
sample_rate: usize,
out_sample_rate: usize,
sr_bin_boundaries: Option<Vec<usize>>,
cond_type: Option<String>,
// cond_dim: Option<usize>,
// cond_out_layer: Option<bool>,
) -> Result<Self> {
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<Tensor> {
let x = self.decoder.forward(z)?;
pub fn decode(&self, z: &Tensor, sr_cond: Option<usize>) -> Result<Tensor> {
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)
}
+9
View File
@@ -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<usize>,
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<bool>,
}
#[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<usize>,
}
#[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<usize>,
pub mean_mode: Option<bool>,
pub cfm_config: CfmConfig,
}
@@ -58,17 +63,21 @@ pub struct AudioVaeConfig {
pub latent_dim: usize,
pub decoder_dim: usize,
pub decoder_rates: Vec<usize>,
pub sr_bin_boundaries: Option<Vec<usize>>,
pub sample_rate: usize,
pub out_sample_rate: Option<usize>,
}
#[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<bool>,
pub encoder_config: VoxCPMEncoderConfig,
pub dit_config: VoxCPMDitConfig,
pub audio_vae_config: Option<AudioVaeConfig>,
+33 -14
View File
@@ -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<HashMap<String, Tensor>>,
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<ChatCompletionResponse> {
let prompt_text = extract_metadata_value::<String>(&mes.metadata, "prompt_text");
let control_instruction =
extract_metadata_value::<String>(&mes.metadata, "control_instruction");
let min_len = extract_metadata_value::<usize>(&mes.metadata, "min_len").unwrap_or(2);
let max_len = extract_metadata_value::<usize>(&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::<f64>(&mes.metadata, "retry_badcase_ratio_threshold")
.unwrap_or(6.0);
let target_text = extract_user_text(&mes)?;
let prompt_wav = extract_audio_url(&mes);
let prompt_wav_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();
+43 -15
View File
@@ -25,7 +25,9 @@ pub struct MiniCPMLongRoPE {
}
impl MiniCPMLongRoPE {
pub fn new(cfg: &VoxMiniCPM4Config, device: &Device, dtype: DType) -> Result<Self> {
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<Tensor> {
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<Tensor> {
let residual = xs.clone();
@@ -232,7 +237,7 @@ pub struct MiniCPMModel {
pub embed_tokens: Option<Embedding>,
layers: Vec<MiniCPMDecoderLayer>,
norm: RmsNorm,
rope_emb: MiniCPMLongRoPE,
rope_emb: Option<MiniCPMLongRoPE>,
}
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(),
)?;
}
+194 -87
View File
@@ -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<Self> {
// 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<Linear>,
stop_proj: Linear,
stop_head: Linear,
device: Device,
@@ -390,17 +411,17 @@ impl VoxCPMModel {
audio_vae: AudioVAE,
) -> Result<Self> {
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<Tensor> {
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::<f32>()? == 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") {
+23
View File
@@ -698,6 +698,29 @@ pub fn bytes_to_human(bytes: u64) -> String {
}
}
pub fn bucketize(input: usize, boundaries: &[usize]) -> Result<usize> {
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::*;
+49
View File
@@ -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(())
}
+23
View File
@@ -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(())
}