refact voxcpm
This commit is contained in:
@@ -23,6 +23,7 @@ pub mod qwen3_reranker;
|
|||||||
pub mod qwen3vl;
|
pub mod qwen3vl;
|
||||||
pub mod rmbg2_0;
|
pub mod rmbg2_0;
|
||||||
pub mod voxcpm;
|
pub mod voxcpm;
|
||||||
|
pub mod voxcpm_refact;
|
||||||
pub mod w2v_bert_2_0;
|
pub mod w2v_bert_2_0;
|
||||||
// pub mod sam3;
|
// pub mod sam3;
|
||||||
pub mod fire_red_vad;
|
pub mod fire_red_vad;
|
||||||
|
|||||||
@@ -265,6 +265,7 @@ impl GenerateModel for VoxCPMGenerate {
|
|||||||
self.voxcpm.clear_kv_cache();
|
self.voxcpm.clear_kv_cache();
|
||||||
})?;
|
})?;
|
||||||
let wav_u8 = get_audio_wav_u8(&audio, self.out_sample_rate as u32)?;
|
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 base64_audio = BASE64_STANDARD.encode(wav_u8);
|
||||||
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
let response = build_audio_completion_response(&base64_audio, &self.model_name);
|
||||||
self.voxcpm.clear_kv_cache();
|
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;
|
use tokenizers::Tokenizer;
|
||||||
|
|
||||||
pub struct SingleChineseTokenizer {
|
pub struct SingleChineseTokenizer {
|
||||||
@@ -62,4 +63,9 @@ impl SingleChineseTokenizer {
|
|||||||
.collect();
|
.collect();
|
||||||
Ok(ids)
|
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)?;
|
let data = BASE64_STANDARD.decode(data)?;
|
||||||
Ok(data)
|
Ok(data)
|
||||||
} else {
|
} 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 std::time::Instant;
|
||||||
|
|
||||||
|
use aha::models::voxcpm_refact::generate::VoxCPMGenerateRefact;
|
||||||
use aha::params::chat::ChatCompletionParameters;
|
use aha::params::chat::ChatCompletionParameters;
|
||||||
use aha::{
|
use aha::{
|
||||||
models::{
|
models::{
|
||||||
@@ -59,7 +60,7 @@ fn voxcpm1_5_use_message_generate() -> Result<()> {
|
|||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn voxcpm1_5_generate() -> Result<()> {
|
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 =
|
let save_dir =
|
||||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get 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 model_path = format!("{}/OpenBMB/VoxCPM1.5/", save_dir);
|
||||||
@@ -71,35 +72,37 @@ fn voxcpm1_5_generate() -> Result<()> {
|
|||||||
|
|
||||||
let i_start = Instant::now();
|
let i_start = Instant::now();
|
||||||
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
// let generate = voxcpm_generate.generate_simple("太阳当空照,花儿对我笑,小鸟说早早早".to_string())?;
|
||||||
let generate = voxcpm_generate.inference(
|
// let generate = voxcpm_generate.inference(
|
||||||
"老大爷我来啦,红红火火恍恍惚惚".to_string(),
|
// "老大爷我来啦,红红火火恍恍惚惚".to_string(),
|
||||||
Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()),
|
// Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()),
|
||||||
Some("file://./assets/audio/voice_01.wav".to_string()),
|
// Some("file://./assets/audio/voice_01.wav".to_string()),
|
||||||
// Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
// // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()),
|
||||||
// Some("file://./assets/audio/voice_05.wav".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,
|
2,
|
||||||
4096,
|
100,
|
||||||
10,
|
10,
|
||||||
2.0,
|
2.0,
|
||||||
// false,
|
false,
|
||||||
6.0,
|
6.0,
|
||||||
)?;
|
)?;
|
||||||
|
|
||||||
// 创建prompt_cache
|
std::thread::sleep(std::time::Duration::from_secs(2));
|
||||||
// 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,
|
|
||||||
// )?;
|
|
||||||
|
|
||||||
let i_duration = i_start.elapsed();
|
let i_duration = i_start.elapsed();
|
||||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||||
@@ -118,3 +121,58 @@ fn voxcpm1_5_tokenizer() -> Result<()> {
|
|||||||
println!("ids: {:?}", ids);
|
println!("ids: {:?}", ids);
|
||||||
Ok(())
|
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