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
+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();