add VoxCPM2
This commit is contained in:
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user