update fmt
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 {
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ fn voxcpm_generate() -> Result<()> {
|
|||||||
100,
|
100,
|
||||||
10,
|
10,
|
||||||
2.0,
|
2.0,
|
||||||
false,
|
// false,
|
||||||
6.0,
|
6.0,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ fn voxcpm1_5_generate() -> Result<()> {
|
|||||||
4096,
|
4096,
|
||||||
10,
|
10,
|
||||||
2.0,
|
2.0,
|
||||||
false,
|
// false,
|
||||||
6.0,
|
6.0,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ fn voxcpm1_5_weight() -> Result<()> {
|
|||||||
dict_to_hashmap.insert(k, v);
|
dict_to_hashmap.insert(k, v);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user