refact voxcpm
This commit is contained in:
@@ -23,6 +23,7 @@ pub mod qwen3_reranker;
|
||||
pub mod qwen3vl;
|
||||
pub mod rmbg2_0;
|
||||
pub mod voxcpm;
|
||||
pub mod voxcpm_refact;
|
||||
pub mod w2v_bert_2_0;
|
||||
// pub mod sam3;
|
||||
pub mod fire_red_vad;
|
||||
|
||||
@@ -265,6 +265,7 @@ impl GenerateModel for VoxCPMGenerate {
|
||||
self.voxcpm.clear_kv_cache();
|
||||
})?;
|
||||
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
||||
// let wave_u8_str = String::from_utf8(wav_u8)?;
|
||||
let base64_audio = BASE64_STANDARD.encode(wav_u8);
|
||||
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
||||
self.voxcpm.clear_kv_cache();
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use anyhow::{Ok, Result, anyhow};
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{Device, Tensor};
|
||||
use tokenizers::Tokenizer;
|
||||
|
||||
pub struct SingleChineseTokenizer {
|
||||
@@ -62,4 +63,9 @@ impl SingleChineseTokenizer {
|
||||
.collect();
|
||||
Ok(ids)
|
||||
}
|
||||
|
||||
pub fn encode_tensor(&self, text: String, device: &Device) -> Result<(Tensor, usize)> {
|
||||
let ids = self.encode(text)?;
|
||||
Ok((Tensor::from_slice(&ids, ids.len(), device)?, ids.len()))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
|
||||
use candle_nn::VarBuilder;
|
||||
use rocket::futures::Stream;
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::{
|
||||
models::{
|
||||
voxcpm::{
|
||||
audio_vae::AudioVAE,
|
||||
config::{AudioVaeConfig, VoxCPMConfig},
|
||||
tokenizer::SingleChineseTokenizer,
|
||||
},
|
||||
voxcpm_refact::{model::VoxCPMModelRefact, processor::VoxCPMProcessor},
|
||||
},
|
||||
utils::{find_type_files, get_device, get_dtype},
|
||||
};
|
||||
|
||||
pub struct VoxCPMGenerateRefact {
|
||||
voxcpm: VoxCPMModelRefact,
|
||||
tokenizer: SingleChineseTokenizer,
|
||||
audio_vae: AudioVAE,
|
||||
processor: VoxCPMProcessor,
|
||||
prompt_cache: Option<HashMap<String, Tensor>>,
|
||||
out_sample_rate: usize,
|
||||
// model_name: String,
|
||||
}
|
||||
|
||||
impl VoxCPMGenerateRefact {
|
||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||
let device = &get_device(device);
|
||||
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")?;
|
||||
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 {
|
||||
dict_to_hashmap.insert(k, v);
|
||||
}
|
||||
}
|
||||
let vb_vae = VarBuilder::from_tensors(dict_to_hashmap, vae_dtype, device);
|
||||
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],
|
||||
sample_rate: 16000,
|
||||
out_sample_rate: None,
|
||||
sr_bin_boundaries: None,
|
||||
},
|
||||
};
|
||||
// let model_name = std::path::Path::new(path)
|
||||
// .file_name()
|
||||
// .and_then(|s| s.to_str())
|
||||
// .unwrap_or("VoxCPM")
|
||||
// .to_string();
|
||||
let audio_vae = AudioVAE::new(
|
||||
vb_vae,
|
||||
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,
|
||||
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 processor = VoxCPMProcessor::new(
|
||||
audio_vae.sample_rate,
|
||||
audio_vae.chunk_size,
|
||||
config.patch_size,
|
||||
device.clone(),
|
||||
);
|
||||
|
||||
let cfg_dtype = config.dtype.as_str();
|
||||
let m_dtype = get_dtype(dtype, cfg_dtype);
|
||||
|
||||
let model_list = find_type_files(path, "bin")?;
|
||||
// voxcpm0.5B模型文件是.bin类型, OpenBMB/VoxCPM1.5模型文件是.safetensors类型
|
||||
let vb_voxcpm = if model_list.is_empty() {
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
unsafe { VarBuilder::from_mmaped_safetensors(&model_list, m_dtype, device)? }
|
||||
} else {
|
||||
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);
|
||||
}
|
||||
}
|
||||
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
|
||||
};
|
||||
let tokenizer = SingleChineseTokenizer::new(path)?;
|
||||
let voxcpm = VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim)?;
|
||||
let out_sample_rate = audio_config
|
||||
.out_sample_rate
|
||||
.unwrap_or(audio_config.sample_rate);
|
||||
Ok(Self {
|
||||
voxcpm,
|
||||
tokenizer,
|
||||
audio_vae,
|
||||
processor,
|
||||
prompt_cache: None,
|
||||
out_sample_rate,
|
||||
// model_name,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn sample_rate(&self) -> usize {
|
||||
self.out_sample_rate
|
||||
}
|
||||
|
||||
pub fn build_prompt_cache(
|
||||
&mut self,
|
||||
prompt_text: String,
|
||||
prompt_wav_path: String,
|
||||
) -> Result<()> {
|
||||
let prompt_cache = self.processor.build_prompt_cache(
|
||||
prompt_text,
|
||||
prompt_wav_path,
|
||||
&self.tokenizer,
|
||||
&self.audio_vae,
|
||||
)?;
|
||||
self.prompt_cache = Some(prompt_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 (text_token, audio_feat, audio_mask) =
|
||||
self.processor
|
||||
.processor_use_cache(target_text, cache, &self.tokenizer)?;
|
||||
let target_text_length = if let Some(mask) = &audio_mask {
|
||||
text_token.dim(1)? - (mask.sum_all()?.to_scalar::<u32>()? as usize)
|
||||
} else {
|
||||
text_token.dim(1)?
|
||||
};
|
||||
let max_len = if retry_badcase {
|
||||
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
|
||||
} else {
|
||||
max_len
|
||||
};
|
||||
self.voxcpm.inference(
|
||||
&text_token,
|
||||
audio_feat.as_ref(),
|
||||
audio_mask.as_ref(),
|
||||
min_len,
|
||||
max_len,
|
||||
inference_timesteps,
|
||||
cfg_value,
|
||||
&self.audio_vae,
|
||||
)?
|
||||
}
|
||||
None => {
|
||||
return Err(anyhow!("need prompt_cache"));
|
||||
}
|
||||
};
|
||||
self.voxcpm.clear_kv_cache();
|
||||
Ok(audio)
|
||||
}
|
||||
|
||||
pub fn generate_stream_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<impl Stream<Item = Result<Tensor, anyhow::Error>>> {
|
||||
match &self.prompt_cache {
|
||||
Some(cache) => {
|
||||
let (text_token, audio_feat, audio_mask) =
|
||||
self.processor
|
||||
.processor_use_cache(target_text, cache, &self.tokenizer)?;
|
||||
let target_text_length = if let Some(mask) = &audio_mask {
|
||||
text_token.dim(1)? - (mask.sum_all()?.to_scalar::<u32>()? as usize)
|
||||
} else {
|
||||
text_token.dim(1)?
|
||||
};
|
||||
let max_len = if retry_badcase {
|
||||
(target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize
|
||||
} else {
|
||||
max_len
|
||||
};
|
||||
self.voxcpm.inference_stream(
|
||||
text_token,
|
||||
audio_feat,
|
||||
audio_mask,
|
||||
min_len,
|
||||
max_len,
|
||||
inference_timesteps,
|
||||
cfg_value,
|
||||
&self.audio_vae,
|
||||
)
|
||||
}
|
||||
None => Err(anyhow!("need prompt_cache")),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
@@ -0,0 +1,469 @@
|
||||
use crate::{
|
||||
models::voxcpm::{
|
||||
audio_vae::AudioVAE,
|
||||
config::VoxCPMConfig,
|
||||
minicpm4::MiniCPMModel,
|
||||
model::{ScalarQuantizationLayer, UnifiedCFM, VoxCPMLocDiT, VoxCPMLocEnc},
|
||||
},
|
||||
utils::tensor_utils::masked_scatter_dim0,
|
||||
};
|
||||
use anyhow::Result;
|
||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||
use candle_nn::{Linear, Module, VarBuilder, linear, linear_no_bias};
|
||||
use rocket::async_stream::stream;
|
||||
use rocket::futures::Stream;
|
||||
|
||||
pub struct VoxCPMModelRefact {
|
||||
config: VoxCPMConfig,
|
||||
patch_size: usize,
|
||||
latent_dim: usize,
|
||||
// audio_start_token: u32,
|
||||
// // audio_end_token: u32,
|
||||
// ref_audio_start_token: u32,
|
||||
// ref_audio_end_token: u32,
|
||||
// chunk_size: usize,
|
||||
// sample_rate: usize,
|
||||
base_lm: MiniCPMModel,
|
||||
residual_lm: MiniCPMModel,
|
||||
feat_encoder: VoxCPMLocEnc,
|
||||
feat_decoder: UnifiedCFM,
|
||||
fsq_layer: ScalarQuantizationLayer,
|
||||
enc_to_lm_proj: Linear,
|
||||
lm_to_dit_proj: Linear,
|
||||
res_to_dit_proj: Linear,
|
||||
fusion_concat_proj: Option<Linear>,
|
||||
stop_proj: Linear,
|
||||
stop_head: Linear,
|
||||
device: Device,
|
||||
dtype: DType,
|
||||
}
|
||||
|
||||
impl VoxCPMModelRefact {
|
||||
pub fn new(vb: VarBuilder, config: VoxCPMConfig, latent_dim: usize) -> Result<Self> {
|
||||
let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
|
||||
let mut residual_lm_config = config.lm_config.clone();
|
||||
residual_lm_config.num_hidden_layers = config.residual_lm_num_layers;
|
||||
residual_lm_config.vocab_size = 0;
|
||||
residual_lm_config.no_rope = config.residual_lm_no_rope;
|
||||
let residual_lm = MiniCPMModel::new(vb.pp("residual_lm"), residual_lm_config)?;
|
||||
let mut encoder_config = config.lm_config.clone();
|
||||
encoder_config.hidden_size = config.encoder_config.hidden_dim;
|
||||
encoder_config.intermediate_size = config.encoder_config.ffn_dim;
|
||||
encoder_config.num_attention_heads = config.encoder_config.num_heads;
|
||||
encoder_config.num_hidden_layers = config.encoder_config.num_layers;
|
||||
encoder_config.kv_channels = config.encoder_config.kv_channels;
|
||||
encoder_config.vocab_size = 0;
|
||||
let feat_encoder =
|
||||
VoxCPMLocEnc::new(vb.pp("feat_encoder"), encoder_config, config.feat_dim)?;
|
||||
|
||||
let mut decoder_config = config.lm_config.clone();
|
||||
decoder_config.hidden_size = config.dit_config.hidden_dim;
|
||||
decoder_config.intermediate_size = config.dit_config.ffn_dim;
|
||||
decoder_config.num_attention_heads = config.dit_config.num_heads;
|
||||
decoder_config.num_hidden_layers = config.dit_config.num_layers;
|
||||
decoder_config.kv_channels = config.dit_config.kv_channels;
|
||||
decoder_config.vocab_size = 0;
|
||||
let estimator = VoxCPMLocDiT::new(
|
||||
vb.pp("feat_decoder.estimator"),
|
||||
decoder_config,
|
||||
config.feat_dim,
|
||||
)?;
|
||||
let feat_decoder = UnifiedCFM::new(
|
||||
config.feat_dim,
|
||||
config.dit_config.cfm_config.clone(),
|
||||
estimator,
|
||||
false,
|
||||
// config.architecture.clone(),
|
||||
)?;
|
||||
let fsq_layer = ScalarQuantizationLayer::new(
|
||||
vb.pp("fsq_layer"),
|
||||
config.lm_config.hidden_size,
|
||||
config.lm_config.hidden_size,
|
||||
config.scalar_quantization_latent_dim,
|
||||
config.scalar_quantization_scale,
|
||||
)?;
|
||||
let enc_to_lm_proj = linear(
|
||||
config.encoder_config.hidden_dim,
|
||||
config.lm_config.hidden_size,
|
||||
vb.pp("enc_to_lm_proj"),
|
||||
)?;
|
||||
let lm_to_dit_proj = linear(
|
||||
config.lm_config.hidden_size,
|
||||
config.dit_config.hidden_dim,
|
||||
vb.pp("lm_to_dit_proj"),
|
||||
)?;
|
||||
let res_to_dit_proj = linear(
|
||||
config.lm_config.hidden_size,
|
||||
config.dit_config.hidden_dim,
|
||||
vb.pp("res_to_dit_proj"),
|
||||
)?;
|
||||
|
||||
let fusion_concat_proj = if config.architecture.to_lowercase().eq("voxcpm2") {
|
||||
Some(linear(
|
||||
config.lm_config.hidden_size * 2,
|
||||
config.lm_config.hidden_size,
|
||||
vb.pp("fusion_concat_proj"),
|
||||
)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let stop_proj = linear(
|
||||
config.lm_config.hidden_size,
|
||||
config.lm_config.hidden_size,
|
||||
vb.pp("stop_proj"),
|
||||
)?;
|
||||
let stop_head = linear_no_bias(config.lm_config.hidden_size, 2, vb.pp("stop_head"))?;
|
||||
|
||||
let patch_size = config.patch_size;
|
||||
Ok(Self {
|
||||
config,
|
||||
patch_size,
|
||||
latent_dim,
|
||||
// audio_start_token: 101,
|
||||
// // audio_end_token: 102,
|
||||
// ref_audio_start_token: 103,
|
||||
// ref_audio_end_token: 104,
|
||||
// chunk_size: audio_vae.chunk_size,
|
||||
// sample_rate: audio_vae.sample_rate,
|
||||
// tokenizer,
|
||||
// audio_vae,
|
||||
base_lm,
|
||||
residual_lm,
|
||||
feat_encoder,
|
||||
feat_decoder,
|
||||
fsq_layer,
|
||||
enc_to_lm_proj,
|
||||
lm_to_dit_proj,
|
||||
res_to_dit_proj,
|
||||
fusion_concat_proj,
|
||||
stop_proj,
|
||||
stop_head,
|
||||
device: vb.device().clone(),
|
||||
dtype: vb.dtype(),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn inference(
|
||||
&mut self,
|
||||
text: &Tensor,
|
||||
audio_feat: Option<&Tensor>,
|
||||
audio_mask: Option<&Tensor>,
|
||||
min_len: usize,
|
||||
max_len: usize,
|
||||
inference_timesteps: usize,
|
||||
cfg_value: f64,
|
||||
audio_vae: &AudioVAE,
|
||||
) -> Result<Tensor> {
|
||||
let (b, t) = text.dims2()?;
|
||||
let scale_emb = if self.config.lm_config.use_mup {
|
||||
self.config.lm_config.scale_emb
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
let text_embed = self
|
||||
.base_lm
|
||||
.embed_tokens
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.forward(text)?
|
||||
.affine(scale_emb as f64, 0.0)?;
|
||||
let (combined_embed, mut prefix_feat_cond, feat_embed) = if let Some(audio_feat) =
|
||||
audio_feat
|
||||
&& let Some(audio_mask) = audio_mask
|
||||
{
|
||||
let audio_feat = audio_feat.to_dtype(self.dtype)?;
|
||||
let audio_t = audio_feat.dim(1)?;
|
||||
let feat_embed = self.feat_encoder.forward(&audio_feat)?; // [b, audio_t, h_feat]
|
||||
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?.squeeze(0)?;
|
||||
let embeds = masked_scatter_dim0(&text_embed, &feat_embed, audio_mask)?;
|
||||
let prefix_feat_cond = audio_feat.i((.., audio_t - 1, ..))?;
|
||||
(embeds, prefix_feat_cond, Some(feat_embed))
|
||||
} else {
|
||||
let prefix_feat_cond = Tensor::zeros(
|
||||
(b, self.patch_size, self.latent_dim),
|
||||
self.dtype,
|
||||
&self.device,
|
||||
)?;
|
||||
(text_embed, prefix_feat_cond, None)
|
||||
};
|
||||
let mut pred_feat_seq = Vec::new();
|
||||
// if feat_mask.i((1, t-1))?.to_scalar::<f32>()? == 0.0 {
|
||||
// // TODO for stream
|
||||
// }
|
||||
let mut position_id = 0;
|
||||
let mut seq_len = t;
|
||||
let enc_outputs = self
|
||||
.base_lm
|
||||
.forward_with_cache(&combined_embed, position_id)?;
|
||||
|
||||
let (mut lm_hidden, input_embeds) = if let Some(_) = audio_feat
|
||||
&& let Some(audio_mask) = audio_mask
|
||||
&& let Some(feat_embed) = feat_embed
|
||||
{
|
||||
let fsq_emb = self.fsq_layer.forward(&enc_outputs)?;
|
||||
let audio_mask_broadcast = audio_mask
|
||||
.unsqueeze(D::Minus1)?
|
||||
.broadcast_as(fsq_emb.shape())?;
|
||||
let enc_outputs = audio_mask_broadcast.where_cond(&fsq_emb, &enc_outputs)?;
|
||||
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
|
||||
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
|
||||
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
|
||||
enc_outputs.add(&feat)?
|
||||
};
|
||||
(lm_hidden, input_embeds)
|
||||
} else {
|
||||
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
|
||||
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
enc_outputs
|
||||
};
|
||||
(lm_hidden, input_embeds)
|
||||
};
|
||||
let residual_enc_outputs = self
|
||||
.residual_lm
|
||||
.forward_with_cache(&input_embeds, position_id)?;
|
||||
let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?;
|
||||
for i in 0..max_len {
|
||||
let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit]
|
||||
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit]
|
||||
// let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
|
||||
let dit_hidden = if self.fusion_concat_proj.is_some() {
|
||||
Tensor::cat(&[&dit_hidden_1, &dit_hidden_2], D::Minus1)?
|
||||
} else {
|
||||
dit_hidden_1.add(&dit_hidden_2)?
|
||||
};
|
||||
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
|
||||
let pred_feat = self
|
||||
.feat_decoder
|
||||
.forward(
|
||||
&dit_hidden,
|
||||
inference_timesteps,
|
||||
self.patch_size,
|
||||
&cond,
|
||||
1.0,
|
||||
cfg_value,
|
||||
1.0,
|
||||
true,
|
||||
)?
|
||||
.transpose(1, 2)?; // [b, p, d]
|
||||
let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c]
|
||||
let curr_embed = self.enc_to_lm_proj.forward(&curr_embed)?;
|
||||
pred_feat_seq.push(pred_feat.unsqueeze(1)?);
|
||||
|
||||
prefix_feat_cond = pred_feat;
|
||||
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
|
||||
let stop_flag = self
|
||||
.stop_head
|
||||
.forward(&stop_flag)?
|
||||
.argmax(D::Minus1)?
|
||||
.i(0)?
|
||||
.to_scalar::<u32>()?;
|
||||
if i > min_len && stop_flag == 1 {
|
||||
break;
|
||||
}
|
||||
position_id += seq_len;
|
||||
seq_len = 1;
|
||||
lm_hidden = self
|
||||
.base_lm
|
||||
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
|
||||
.squeeze(1)?;
|
||||
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
|
||||
let curr_residual_input = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let curr_embed = curr_embed.i((.., 0, ..))?;
|
||||
let concat = Tensor::cat(&[&lm_hidden, &curr_embed], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
lm_hidden.add(&curr_embed.i((.., 0, ..))?)?
|
||||
};
|
||||
residual_hidden = self
|
||||
.residual_lm
|
||||
.forward_with_cache(&curr_residual_input, position_id)?
|
||||
.squeeze(1)?;
|
||||
}
|
||||
let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d)
|
||||
let (b, _, _, d) = pred_seq.dims4()?;
|
||||
let feat_pred = pred_seq
|
||||
.permute((0, 3, 1, 2))?
|
||||
.reshape((b, d, ()))?
|
||||
.contiguous()?;
|
||||
self.clear_kv_cache();
|
||||
|
||||
let decode_audio = audio_vae
|
||||
.decode(&feat_pred.to_dtype(DType::F32)?, None)?
|
||||
.squeeze(1)?;
|
||||
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640 - 640;
|
||||
let decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?;
|
||||
Ok(decode_audio)
|
||||
}
|
||||
|
||||
pub fn inference_stream(
|
||||
&mut self,
|
||||
text: Tensor,
|
||||
audio_feat: Option<Tensor>,
|
||||
audio_mask: Option<Tensor>,
|
||||
min_len: usize,
|
||||
max_len: usize,
|
||||
inference_timesteps: usize,
|
||||
cfg_value: f64,
|
||||
audio_vae: &AudioVAE,
|
||||
) -> Result<impl Stream<Item = Result<Tensor, anyhow::Error>>> {
|
||||
let (b, t) = text.dims2()?;
|
||||
let scale_emb = if self.config.lm_config.use_mup {
|
||||
self.config.lm_config.scale_emb
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
let text_embed = self
|
||||
.base_lm
|
||||
.embed_tokens
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.forward(&text)?
|
||||
.affine(scale_emb as f64, 0.0)?;
|
||||
let (combined_embed, mut prefix_feat_cond, feat_embed) = if let Some(audio_feat) =
|
||||
&audio_feat
|
||||
&& let Some(audio_mask) = &audio_mask
|
||||
{
|
||||
let audio_feat = audio_feat.to_dtype(self.dtype)?;
|
||||
let audio_t = audio_feat.dim(1)?;
|
||||
let feat_embed = self.feat_encoder.forward(&audio_feat)?; // [b, audio_t, h_feat]
|
||||
let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?.squeeze(0)?;
|
||||
let embeds = masked_scatter_dim0(&text_embed, &feat_embed, audio_mask)?;
|
||||
let prefix_feat_cond = audio_feat.i((.., audio_t - 1, ..))?;
|
||||
(embeds, prefix_feat_cond, Some(feat_embed))
|
||||
} else {
|
||||
let prefix_feat_cond = Tensor::zeros(
|
||||
(b, self.patch_size, self.latent_dim),
|
||||
self.dtype,
|
||||
&self.device,
|
||||
)?;
|
||||
(text_embed, prefix_feat_cond, None)
|
||||
};
|
||||
// let mut pred_feat_seq = Vec::new();
|
||||
// if feat_mask.i((1, t-1))?.to_scalar::<f32>()? == 0.0 {
|
||||
// // TODO for stream
|
||||
// }
|
||||
let mut position_id = 0;
|
||||
let mut seq_len = t;
|
||||
let enc_outputs = self
|
||||
.base_lm
|
||||
.forward_with_cache(&combined_embed, position_id)?;
|
||||
|
||||
let (mut lm_hidden, input_embeds) = if let Some(_) = &audio_feat
|
||||
&& let Some(audio_mask) = &audio_mask
|
||||
&& let Some(feat_embed) = feat_embed
|
||||
{
|
||||
let fsq_emb = self.fsq_layer.forward(&enc_outputs)?;
|
||||
let audio_mask_broadcast = audio_mask
|
||||
.unsqueeze(D::Minus1)?
|
||||
.broadcast_as(fsq_emb.shape())?;
|
||||
let enc_outputs = audio_mask_broadcast.where_cond(&fsq_emb, &enc_outputs)?;
|
||||
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
|
||||
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
|
||||
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let feat = masked_scatter_dim0(&feat, &feat_embed, audio_mask)?;
|
||||
enc_outputs.add(&feat)?
|
||||
};
|
||||
(lm_hidden, input_embeds)
|
||||
} else {
|
||||
let lm_hidden = enc_outputs.i((.., t - 1, ..))?;
|
||||
let input_embeds = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let feat = enc_outputs.zeros_like()?;
|
||||
let concat = Tensor::cat(&[&enc_outputs, &feat], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
enc_outputs
|
||||
};
|
||||
(lm_hidden, input_embeds)
|
||||
};
|
||||
let residual_enc_outputs = self
|
||||
.residual_lm
|
||||
.forward_with_cache(&input_embeds, position_id)?;
|
||||
let mut residual_hidden = residual_enc_outputs.i((.., t - 1, ..))?;
|
||||
let stream = stream! {
|
||||
for i in 0..max_len {
|
||||
let dit_hidden_1 = self.lm_to_dit_proj.forward(&lm_hidden)?; // [b, h_dit]
|
||||
let dit_hidden_2 = self.res_to_dit_proj.forward(&residual_hidden)?; // [b, h_dit]
|
||||
// let dit_hidden = dit_hidden_1.add(&dit_hidden_2)?;
|
||||
let dit_hidden = if self.fusion_concat_proj.is_some() {
|
||||
Tensor::cat(&[&dit_hidden_1, &dit_hidden_2], D::Minus1)?
|
||||
} else {
|
||||
dit_hidden_1.add(&dit_hidden_2)?
|
||||
};
|
||||
let cond = prefix_feat_cond.transpose(1, 2)?.contiguous()?;
|
||||
let pred_feat = self
|
||||
.feat_decoder
|
||||
.forward(
|
||||
&dit_hidden,
|
||||
inference_timesteps,
|
||||
self.patch_size,
|
||||
&cond,
|
||||
1.0,
|
||||
cfg_value,
|
||||
1.0,
|
||||
true,
|
||||
)?
|
||||
.transpose(1, 2)?; // [b, p, d]
|
||||
let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c]
|
||||
let curr_embed = self.enc_to_lm_proj.forward(&curr_embed)?;
|
||||
let single_feat_pred = pred_feat.permute((0, 2, 1))?.contiguous()?;
|
||||
let decode_audio = audio_vae
|
||||
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
|
||||
.squeeze(1)?;
|
||||
yield Ok(decode_audio);
|
||||
prefix_feat_cond = pred_feat;
|
||||
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
|
||||
let stop_flag = self
|
||||
.stop_head
|
||||
.forward(&stop_flag)?
|
||||
.argmax(D::Minus1)?
|
||||
.i(0)?
|
||||
.to_scalar::<u32>()?;
|
||||
if i > min_len && stop_flag == 1 {
|
||||
break;
|
||||
}
|
||||
position_id += seq_len;
|
||||
seq_len = 1;
|
||||
lm_hidden = self
|
||||
.base_lm
|
||||
.forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)?
|
||||
.squeeze(1)?;
|
||||
lm_hidden = self.fsq_layer.forward(&lm_hidden)?;
|
||||
let curr_residual_input = if let Some(fusion) = &self.fusion_concat_proj {
|
||||
let curr_embed = curr_embed.i((.., 0, ..))?;
|
||||
let concat = Tensor::cat(&[&lm_hidden, &curr_embed], D::Minus1)?;
|
||||
fusion.forward(&concat)?
|
||||
} else {
|
||||
lm_hidden.add(&curr_embed.i((.., 0, ..))?)?
|
||||
};
|
||||
residual_hidden = self
|
||||
.residual_lm
|
||||
.forward_with_cache(&curr_residual_input, position_id)?
|
||||
.squeeze(1)?;
|
||||
|
||||
}
|
||||
self.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
}
|
||||
|
||||
pub fn clear_kv_cache(&mut self) {
|
||||
self.base_lm.clear_kv_cache();
|
||||
self.residual_lm.clear_kv_cache();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::{
|
||||
models::voxcpm::{audio_vae::AudioVAE, tokenizer::SingleChineseTokenizer},
|
||||
utils::audio_utils::load_audio_with_resample,
|
||||
};
|
||||
use anyhow::Result;
|
||||
use candle_core::{D, Device, IndexOp, Tensor};
|
||||
|
||||
pub struct VoxCPMProcessor {
|
||||
sample_rate: usize,
|
||||
chunk_size: usize,
|
||||
patch_size: usize,
|
||||
audio_start_token: u32,
|
||||
ref_audio_start_token: u32,
|
||||
ref_audio_end_token: u32,
|
||||
device: Device,
|
||||
}
|
||||
|
||||
impl VoxCPMProcessor {
|
||||
pub fn new(sample_rate: usize, chunk_size: usize, patch_size: usize, device: Device) -> Self {
|
||||
Self {
|
||||
sample_rate,
|
||||
chunk_size,
|
||||
patch_size,
|
||||
audio_start_token: 101,
|
||||
ref_audio_start_token: 103,
|
||||
ref_audio_end_token: 104,
|
||||
device,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_prompt_cache(
|
||||
&mut self,
|
||||
prompt_text: String,
|
||||
prompt_wav_path: String,
|
||||
tokenizer: &SingleChineseTokenizer,
|
||||
audio_vae: &AudioVAE,
|
||||
) -> Result<HashMap<String, Tensor>> {
|
||||
let (text_token, _) = tokenizer.encode_tensor(prompt_text, &self.device)?;
|
||||
let mut audio =
|
||||
load_audio_with_resample(&prompt_wav_path, &self.device, Some(self.sample_rate))?;
|
||||
let patch_len = self.patch_size * self.chunk_size;
|
||||
if audio.dim(1)? % patch_len != 0 {
|
||||
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
|
||||
}
|
||||
let audio_feat = audio_vae.encode(&audio, Some(self.sample_rate))?;
|
||||
let audio_feat = audio_feat
|
||||
.reshape((audio_vae.latent_dim, (), self.patch_size))?
|
||||
.permute((1, 2, 0))?;
|
||||
let dim0 = audio_feat.dim(0)? - 1;
|
||||
let audio_feat = audio_feat.i(..dim0)?;
|
||||
let mut hashmap = HashMap::new();
|
||||
hashmap.insert("text_token".to_string(), text_token);
|
||||
hashmap.insert("audio_feat".to_string(), audio_feat);
|
||||
Ok(hashmap)
|
||||
}
|
||||
|
||||
pub fn processor(
|
||||
&self,
|
||||
target_text: String,
|
||||
prompt_text: Option<String>,
|
||||
prompt_wav_path: Option<String>,
|
||||
tokenizer: &SingleChineseTokenizer,
|
||||
audio_vae: &AudioVAE,
|
||||
) -> Result<(Tensor, Option<Tensor>, Option<Tensor>)> {
|
||||
let text = if let Some(prompt_text) = &prompt_text {
|
||||
prompt_text.clone() + &target_text
|
||||
} else {
|
||||
target_text
|
||||
};
|
||||
let (text_token, _) = tokenizer.encode_tensor(text, &self.device)?;
|
||||
let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
|
||||
let mut text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
|
||||
|
||||
let (audio_feat, audio_mask) = if let Some(path) = prompt_wav_path {
|
||||
let mut audio = load_audio_with_resample(&path, &self.device, Some(self.sample_rate))?;
|
||||
let patch_len = self.patch_size * self.chunk_size;
|
||||
if audio.dim(1)? % patch_len != 0 {
|
||||
audio =
|
||||
audio.pad_with_zeros(D::Minus1, patch_len - audio.dim(1)? % patch_len, 0)?;
|
||||
}
|
||||
let audio_feat = audio_vae.encode(&audio, Some(self.sample_rate))?;
|
||||
let audio_feat = audio_feat
|
||||
.reshape((audio_vae.latent_dim, (), self.patch_size))?
|
||||
.permute((1, 2, 0))?;
|
||||
let text_length = text_token.dim(0)?;
|
||||
let audio_length = audio_feat.dim(0)?;
|
||||
let audio_mask = if prompt_text.is_some() {
|
||||
let text_pad_token =
|
||||
Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
|
||||
text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
|
||||
let mask = Tensor::cat(
|
||||
&[
|
||||
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
|
||||
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
|
||||
],
|
||||
D::Minus1,
|
||||
)?
|
||||
.unsqueeze(0)?;
|
||||
Some(mask)
|
||||
} else {
|
||||
let ref_start = Tensor::new(vec![self.ref_audio_start_token], &self.device)?;
|
||||
let ref_end = Tensor::new(vec![self.ref_audio_end_token], &self.device)?;
|
||||
let ref_token = Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
|
||||
text_token = Tensor::cat(&[&ref_start, &ref_token, &ref_end, &text_token], 0)?;
|
||||
let mask = Tensor::cat(
|
||||
&[
|
||||
Tensor::new(vec![0u32], &self.device)?,
|
||||
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
|
||||
Tensor::new(vec![0u32], &self.device)?,
|
||||
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
|
||||
],
|
||||
D::Minus1,
|
||||
)?
|
||||
.unsqueeze(0)?;
|
||||
Some(mask)
|
||||
};
|
||||
let audio_feat = audio_feat.unsqueeze(0)?;
|
||||
(Some(audio_feat), audio_mask)
|
||||
} else {
|
||||
(None, None)
|
||||
};
|
||||
let text_token = text_token.unsqueeze(0)?;
|
||||
Ok((text_token, audio_feat, audio_mask))
|
||||
}
|
||||
|
||||
pub fn processor_use_cache(
|
||||
&self,
|
||||
target_text: String,
|
||||
prompt_cache: &HashMap<String, Tensor>,
|
||||
tokenizer: &SingleChineseTokenizer,
|
||||
) -> Result<(Tensor, Option<Tensor>, Option<Tensor>)> {
|
||||
let (target_text_token, _) = tokenizer.encode_tensor(target_text, &self.device)?;
|
||||
let text_token = match prompt_cache.get("text_token") {
|
||||
Some(token) => Tensor::cat(&[token, &target_text_token], 0)?,
|
||||
None => target_text_token,
|
||||
};
|
||||
let audio_start = Tensor::new(vec![self.audio_start_token], &self.device)?;
|
||||
let mut text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
|
||||
let text_length = text_token.dim(0)?;
|
||||
let (audio_length, audio_feat) = match prompt_cache.get("audio_feat") {
|
||||
Some(feat) => (feat.dim(0)?, Some(feat.clone().unsqueeze(0)?)),
|
||||
None => (0, None),
|
||||
};
|
||||
let audio_mask = if audio_length > 0 {
|
||||
let text_pad_token =
|
||||
Tensor::zeros(audio_length, candle_core::DType::U32, &self.device)?;
|
||||
text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?;
|
||||
let mask = Tensor::cat(
|
||||
&[
|
||||
Tensor::zeros(text_length, candle_core::DType::U32, &self.device)?,
|
||||
Tensor::ones(audio_length, candle_core::DType::U32, &self.device)?,
|
||||
],
|
||||
D::Minus1,
|
||||
)?
|
||||
.unsqueeze(0)?;
|
||||
Some(mask)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let text_token = text_token.unsqueeze(0)?;
|
||||
Ok((text_token, audio_feat, audio_mask))
|
||||
}
|
||||
}
|
||||
@@ -366,7 +366,15 @@ pub fn get_audio_bytes_vec(path_str: &str) -> Result<Vec<u8>> {
|
||||
let data = BASE64_STANDARD.decode(data)?;
|
||||
Ok(data)
|
||||
} else {
|
||||
Err(anyhow::anyhow!("get audio path error {}", path_str))
|
||||
let wave_u8 = path_str.as_bytes();
|
||||
match get_audio_format_from_bytes(wave_u8) {
|
||||
Ok(_) => Ok(wave_u8.to_vec()),
|
||||
Err(e) => Err(anyhow::anyhow!(
|
||||
"get audio path error {}, et_audio_format error: {}",
|
||||
path_str,
|
||||
e
|
||||
)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+82
-24
@@ -1,5 +1,6 @@
|
||||
use std::time::Instant;
|
||||
|
||||
use aha::models::voxcpm_refact::generate::VoxCPMGenerateRefact;
|
||||
use aha::params::chat::ChatCompletionParameters;
|
||||
use aha::{
|
||||
models::{
|
||||
@@ -59,7 +60,7 @@ fn voxcpm1_5_use_message_generate() -> Result<()> {
|
||||
|
||||
#[test]
|
||||
fn voxcpm1_5_generate() -> Result<()> {
|
||||
// RUST_BACKTRACE=1 cargo test -F cuda voxcpm1_5_generate -r -- --nocapture
|
||||
// RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm1_5 voxcpm1_5_generate -r -- --nocapture
|
||||
let save_dir =
|
||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||
let model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir);
|
||||
@@ -71,35 +72,37 @@ fn voxcpm1_5_generate() -> Result<()> {
|
||||
|
||||
let i_start = Instant::now();
|
||||
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
||||
let generate = voxcpm_generate.inference(
|
||||
"老大爷我来啦,红红火火恍恍惚惚".to_string(),
|
||||
Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()),
|
||||
Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||
// Some("file://./assets/audio/voice_05.wav".to_string()),
|
||||
// let generate = voxcpm_generate.inference(
|
||||
// "老大爷我来啦,红红火火恍恍惚惚".to_string(),
|
||||
// Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()),
|
||||
// Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||
// // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||
// // Some("file://./assets/audio/voice_05.wav".to_string()),
|
||||
// 2,
|
||||
// 4096,
|
||||
// 10,
|
||||
// 2.0,
|
||||
// // false,
|
||||
// 6.0,
|
||||
// )?;
|
||||
|
||||
// 创建prompt_cache
|
||||
voxcpm_generate.build_prompt_cache(
|
||||
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||
"file://./assets/audio/voice_01.wav".to_string(),
|
||||
)?;
|
||||
// 使用prompt_cache生成语音
|
||||
let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
||||
2,
|
||||
4096,
|
||||
100,
|
||||
10,
|
||||
2.0,
|
||||
// false,
|
||||
false,
|
||||
6.0,
|
||||
)?;
|
||||
|
||||
// 创建prompt_cache
|
||||
// let _ = voxcpm_generate.build_prompt_cache(
|
||||
// "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||
// "file://./assets/audio/voice_01.wav".to_string(),
|
||||
// )?;
|
||||
// // 使用prompt_cache生成语音
|
||||
// let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||
// "太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
||||
// 2,
|
||||
// 100,
|
||||
// 10,
|
||||
// 2.0,
|
||||
// false,
|
||||
// 6.0,
|
||||
// )?;
|
||||
std::thread::sleep(std::time::Duration::from_secs(2));
|
||||
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
@@ -118,3 +121,58 @@ fn voxcpm1_5_tokenizer() -> Result<()> {
|
||||
println!("ids: {:?}", ids);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn voxcpm_refact_generate() -> Result<()> {
|
||||
// RUST_BACKTRACE=1 cargo test -F cuda --test test_voxcpm1_5 voxcpm_refact_generate -r -- --nocapture
|
||||
let save_dir =
|
||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||
let model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let mut voxcpm_generate = VoxCPMGenerateRefact::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
// let i_start = Instant::now();
|
||||
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
||||
// let generate = voxcpm_generate.inference(
|
||||
// "VoxCPM is an innovative end-to-end TTS model from ModelBest, designed to generate highly realistic speech.".to_string(),
|
||||
// Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()),
|
||||
// Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||
// // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||
// // Some("file://./assets/audio/voice_05.wav".to_string()),
|
||||
// 2,
|
||||
// 100,
|
||||
// 10,
|
||||
// 2.0,
|
||||
// // false,
|
||||
// 6.0,
|
||||
// )?;
|
||||
|
||||
// 创建prompt_cache
|
||||
voxcpm_generate.build_prompt_cache(
|
||||
"啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(),
|
||||
"file://./assets/audio/voice_01.wav".to_string(),
|
||||
)?;
|
||||
// 使用prompt_cache生成语音
|
||||
let i_start = Instant::now();
|
||||
let generate = voxcpm_generate.generate_use_prompt_cache(
|
||||
"太阳当空照,花儿对我笑,小鸟说早早早".to_string(),
|
||||
2,
|
||||
100,
|
||||
10,
|
||||
2.0,
|
||||
false,
|
||||
6.0,
|
||||
)?;
|
||||
std::thread::sleep(std::time::Duration::from_secs(2));
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
save_wav(
|
||||
&generate,
|
||||
"voxcpm.wav",
|
||||
voxcpm_generate.sample_rate() as u32,
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user