Files
aha/src/models/voxcpm/generate.rs
T

276 lines
9.2 KiB
Rust
Raw Normal View History

2025-10-10 20:36:52 +08:00
use std::collections::HashMap;
2026-03-31 12:30:12 +08:00
use crate::params::chat::{
2025-12-25 20:25:52 +08:00
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
2025-10-15 21:03:49 +08:00
use anyhow::{Ok, Result};
2025-12-25 20:25:52 +08:00
use base64::{Engine, prelude::BASE64_STANDARD};
2025-10-15 21:03:49 +08:00
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
2025-12-25 20:25:52 +08:00
use rocket::futures::{Stream, stream};
2025-10-15 21:03:49 +08:00
2025-10-10 20:36:52 +08:00
use crate::{
2025-12-25 20:25:52 +08:00
models::{
GenerateModel,
voxcpm::{
audio_vae::AudioVAE,
config::{AudioVaeConfig, VoxCPMConfig},
model::VoxCPMModel,
tokenizer::SingleChineseTokenizer,
},
},
utils::{
audio_utils::{extract_audio_url, get_audio_wav_u8},
build_audio_completion_response, extract_metadata_value, extract_user_text,
find_type_files, get_device, get_dtype,
2025-10-10 20:36:52 +08:00
},
};
pub struct VoxCPMGenerate {
voxcpm: VoxCPMModel,
prompt_cache: Option<HashMap<String, Tensor>>,
2025-12-25 20:25:52 +08:00
sample_rate: usize,
model_name: String,
2025-10-10 20:36:52 +08:00
}
impl VoxCPMGenerate {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let device = &get_device(device);
2025-12-11 18:33:35 +08:00
let config_path = path.to_string() + "/config.json";
let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
2025-10-10 20:36:52 +08:00
let model_list = find_type_files(path, "pth")?;
2025-10-11 11:06:57 +08:00
// println!(" pth model_list: {:?}", model_list);
2025-10-10 20:36:52 +08:00
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);
}
}
2025-10-15 21:03:49 +08:00
let vb_vae = VarBuilder::from_tensors(dict_to_hashmap, vae_dtype, device);
2025-12-11 18:33:35 +08:00
let audio_config = match config.audio_vae_config.clone() {
Some(config) => config,
None => AudioVaeConfig {
encoder_dim: 128,
encoder_rates: vec![2, 5, 8, 8],
latent_dim: 64,
decoder_dim: 1536,
decoder_rates: vec![8, 8, 5, 2],
2025-12-11 23:30:07 +08:00
sample_rate: 16000,
},
2025-12-11 18:33:35 +08:00
};
2026-03-30 20:33:48 +08:00
let model_name = std::path::Path::new(path)
2026-03-30 20:36:57 +08:00
.file_name()
2026-03-30 20:33:48 +08:00
.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()
// };
2025-10-10 20:36:52 +08:00
let audio_vae = AudioVAE::new(
vb_vae,
2025-12-11 18:33:35 +08:00
audio_config.encoder_dim,
audio_config.encoder_rates.clone(),
Some(audio_config.latent_dim),
audio_config.decoder_dim,
audio_config.decoder_rates.clone(),
audio_config.sample_rate,
2025-10-10 20:36:52 +08:00
)?;
2025-10-11 14:16:49 +08:00
let cfg_dtype = config.dtype.as_str();
2025-10-15 21:03:49 +08:00
let m_dtype = get_dtype(dtype, cfg_dtype);
let model_list = find_type_files(path, "bin")?;
// voxcpm0.5B模型文件是.bin类型, voxcpm1.5模型文件是.safetensors类型
let vb_voxcpm = if model_list.is_empty() {
let model_list = find_type_files(path, "safetensors")?;
2025-12-11 23:30:07 +08:00
unsafe { VarBuilder::from_mmaped_safetensors(&model_list, m_dtype, device)? }
} else {
2025-12-11 23:30:07 +08:00
dict_to_hashmap = HashMap::new();
let cfg_dtype = config.dtype.as_str();
let m_dtype = get_dtype(dtype, cfg_dtype);
for m in model_list {
let dict = read_all_with_key(m, Some("state_dict"))?;
for (k, v) in dict {
// println!("key: {}, tensor shape: {:?}", k, v);
dict_to_hashmap.insert(k, v);
}
2025-10-10 20:36:52 +08:00
}
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
};
2025-10-10 20:36:52 +08:00
let tokenizer = SingleChineseTokenizer::new(path)?;
let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?;
Ok(Self {
voxcpm,
prompt_cache: None,
2025-12-25 20:25:52 +08:00
sample_rate: audio_config.sample_rate,
model_name,
2025-10-10 20:36:52 +08:00
})
}
pub fn build_prompt_cache(
&mut self,
prompt_text: String,
prompt_wav_path: String,
) -> Result<()> {
let cache = self
.voxcpm
.build_prompt_cache(prompt_text, prompt_wav_path)?;
self.prompt_cache = Some(cache);
Ok(())
}
pub fn generate_use_prompt_cache(
&mut self,
target_text: String,
min_len: usize,
max_len: usize,
inference_timesteps: usize,
cfg_value: f64,
retry_badcase: bool,
retry_badcase_ratio_threshold: f64,
) -> Result<Tensor> {
let audio = match &self.prompt_cache {
Some(cache) => {
let prompt_cache = cache.clone();
self.voxcpm.generate_with_prompt_cache(
target_text,
prompt_cache,
min_len,
max_len,
inference_timesteps,
cfg_value,
retry_badcase,
retry_badcase_ratio_threshold,
)?
}
None => self.generate_simple(target_text)?,
};
2026-01-15 21:57:12 +08:00
self.voxcpm.clear_kv_cache();
2025-10-10 20:36:52 +08:00
Ok(audio)
}
pub fn generate_with_prompt_simple(
&mut self,
target_text: String,
prompt_text: Option<String>,
prompt_wav_path: Option<String>,
) -> Result<Tensor> {
2025-12-25 20:25:52 +08:00
let audio = self.inference(
2025-10-10 20:36:52 +08:00
target_text,
prompt_text,
prompt_wav_path,
2,
1000,
10,
2.0,
2025-12-11 23:30:07 +08:00
// false,
2025-10-10 20:36:52 +08:00
6.0,
)?;
Ok(audio)
}
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
2025-12-11 23:30:07 +08:00
// let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?;
2025-12-25 20:25:52 +08:00
let audio = self.inference(target_text, None, None, 2, 100, 10, 2.0, 6.0)?;
2025-10-10 20:36:52 +08:00
Ok(audio)
}
2025-12-25 20:25:52 +08:00
pub fn inference(
2025-10-10 20:36:52 +08:00
&mut self,
target_text: String,
prompt_text: Option<String>,
prompt_wav_path: Option<String>,
min_len: usize,
max_len: usize,
inference_timesteps: usize,
cfg_value: f64,
2025-12-11 23:30:07 +08:00
// retry_badcase: bool,
2025-10-10 20:36:52 +08:00
retry_badcase_ratio_threshold: f64,
) -> Result<Tensor> {
let audio = self.voxcpm.generate(
target_text,
prompt_text,
prompt_wav_path,
min_len,
max_len,
inference_timesteps,
cfg_value,
2025-12-11 23:30:07 +08:00
// retry_badcase,
2025-10-10 20:36:52 +08:00
retry_badcase_ratio_threshold,
)?;
2026-01-15 21:57:12 +08:00
self.voxcpm.clear_kv_cache();
2025-10-10 20:36:52 +08:00
Ok(audio)
}
2026-01-08 00:04:24 +08:00
pub fn sample_rate(&self) -> usize {
self.sample_rate
}
2025-10-10 20:36:52 +08:00
}
2025-12-25 20:25:52 +08:00
impl GenerateModel for VoxCPMGenerate {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let prompt_text = extract_metadata_value::<String>(&mes.metadata, "prompt_text");
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 =
extract_metadata_value::<usize>(&mes.metadata, "inference_timesteps").unwrap_or(10);
let cfg_value = extract_metadata_value::<f64>(&mes.metadata, "cfg_value").unwrap_or(2.0);
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)?;
2025-12-31 17:39:25 +08:00
let prompt_wav = extract_audio_url(&mes);
2025-12-25 20:25:52 +08:00
let prompt_wav_path = if !prompt_wav.is_empty() {
Some(prompt_wav[0].clone())
} else {
None
};
2026-01-15 21:57:12 +08:00
let audio = self
.voxcpm
.generate(
target_text,
prompt_text,
prompt_wav_path,
min_len,
max_len,
inference_timesteps,
cfg_value,
retry_badcase_ratio_threshold,
)
.inspect_err(|_| {
self.voxcpm.clear_kv_cache();
})?;
2025-12-25 20:25:52 +08:00
let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate as u32)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name);
2026-01-15 21:57:12 +08:00
self.voxcpm.clear_kv_cache();
2025-12-25 20:25:52 +08:00
Ok(response)
}
#[allow(unused_variables)]
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ Send
+ Unpin
+ '_,
>,
> {
let error_stream = stream::once(async {
Err(anyhow::anyhow!(format!(
"{} model not support stream",
self.model_name
))) as Result<ChatCompletionChunkResponse, anyhow::Error>
});
Ok(Box::new(Box::pin(error_stream)))
}
}