update fmt

This commit is contained in:
jhqxxx
2025-12-11 23:30:07 +08:00
parent d282a0cf40
commit 855d0cf975
6 changed files with 18 additions and 14 deletions
+2 -1
View File
@@ -444,7 +444,8 @@ impl CausalDecoder {
} }
let idx = rates.len() + 2; let idx = rates.len() + 2;
let model_minus_2 = Snake1d::new(vb_model.pp(idx), output_dim)?; 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 model_minus_1 =
WNCausalConv1d::new(vb_model.pp(idx + 1), output_dim, d_out, 7, 1, 3, 1, 1)?;
Ok(Self { Ok(Self {
model0, model0,
model1, model1,
+12 -9
View File
@@ -6,7 +6,9 @@ use candle_nn::VarBuilder;
use crate::{ use crate::{
models::voxcpm::{ models::voxcpm::{
audio_vae::AudioVAE, config::{AudioVaeConfig, VoxCPMConfig}, model::VoxCPMModel, audio_vae::AudioVAE,
config::{AudioVaeConfig, VoxCPMConfig},
model::VoxCPMModel,
tokenizer::SingleChineseTokenizer, tokenizer::SingleChineseTokenizer,
}, },
utils::{find_type_files, get_device, get_dtype}, utils::{find_type_files, get_device, get_dtype},
@@ -43,8 +45,8 @@ impl VoxCPMGenerate {
latent_dim: 64, latent_dim: 64,
decoder_dim: 1536, decoder_dim: 1536,
decoder_rates: vec![8, 8, 5, 2], decoder_rates: vec![8, 8, 5, 2],
sample_rate: 16000 sample_rate: 16000,
} },
}; };
let audio_vae = AudioVAE::new( let audio_vae = AudioVAE::new(
vb_vae, vb_vae,
@@ -63,9 +65,9 @@ impl VoxCPMGenerate {
// voxcpm0.5B模型文件是.bin类型, voxcpm1.5模型文件是.safetensors类型 // voxcpm0.5B模型文件是.bin类型, voxcpm1.5模型文件是.safetensors类型
let vb_voxcpm = if model_list.is_empty() { let vb_voxcpm = if model_list.is_empty() {
let model_list = find_type_files(path, "safetensors")?; let model_list = find_type_files(path, "safetensors")?;
unsafe { VarBuilder::from_mmaped_safetensors(&model_list, m_dtype, &device)? } unsafe { VarBuilder::from_mmaped_safetensors(&model_list, m_dtype, device)? }
} else { } else {
dict_to_hashmap = HashMap::new(); dict_to_hashmap = HashMap::new();
let cfg_dtype = config.dtype.as_str(); let cfg_dtype = config.dtype.as_str();
let m_dtype = get_dtype(dtype, cfg_dtype); let m_dtype = get_dtype(dtype, cfg_dtype);
for m in model_list { for m in model_list {
@@ -141,13 +143,14 @@ impl VoxCPMGenerate {
1000, 1000,
10, 10,
2.0, 2.0,
false, // false,
6.0, 6.0,
)?; )?;
Ok(audio) Ok(audio)
} }
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> { pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?; // let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?;
let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, 6.0)?;
Ok(audio) Ok(audio)
} }
pub fn generate( pub fn generate(
@@ -159,7 +162,7 @@ impl VoxCPMGenerate {
max_len: usize, max_len: usize,
inference_timesteps: usize, inference_timesteps: usize,
cfg_value: f64, cfg_value: f64,
retry_badcase: bool, // retry_badcase: bool,
retry_badcase_ratio_threshold: f64, retry_badcase_ratio_threshold: f64,
) -> Result<Tensor> { ) -> Result<Tensor> {
let audio = self.voxcpm.generate( let audio = self.voxcpm.generate(
@@ -170,7 +173,7 @@ impl VoxCPMGenerate {
max_len, max_len,
inference_timesteps, inference_timesteps,
cfg_value, cfg_value,
retry_badcase, // retry_badcase,
retry_badcase_ratio_threshold, retry_badcase_ratio_threshold,
)?; )?;
Ok(audio) Ok(audio)
+1 -1
View File
@@ -486,7 +486,7 @@ impl VoxCPMModel {
max_len: usize, max_len: usize,
inference_timesteps: usize, inference_timesteps: usize,
cfg_value: f64, cfg_value: f64,
retry_badcase: bool, // retry_badcase: bool,
retry_badcase_ratio_threshold: f64, retry_badcase_ratio_threshold: f64,
) -> Result<Tensor> { ) -> Result<Tensor> {
let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path { let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path {
+1 -1
View File
@@ -28,7 +28,7 @@ fn voxcpm_generate() -> Result<()> {
100, 100,
10, 10,
2.0, 2.0,
false, // false,
6.0, 6.0,
)?; )?;
+1 -1
View File
@@ -28,7 +28,7 @@ fn voxcpm1_5_generate() -> Result<()> {
4096, 4096,
10, 10,
2.0, 2.0,
false, // false,
6.0, 6.0,
)?; )?;
+1 -1
View File
@@ -62,7 +62,7 @@ fn voxcpm1_5_weight() -> Result<()> {
dict_to_hashmap.insert(k, v); dict_to_hashmap.insert(k, v);
} }
} }
Ok(()) Ok(())
} }