diff --git a/src/models/mod.rs b/src/models/mod.rs index 8e3c440..6cb658c 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -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; diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 68cd7c3..e9ebbfb 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -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(); diff --git a/src/models/voxcpm/tokenizer.rs b/src/models/voxcpm/tokenizer.rs index 1e1e5a0..2b8ee27 100644 --- a/src/models/voxcpm/tokenizer.rs +++ b/src/models/voxcpm/tokenizer.rs @@ -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())) + } } diff --git a/src/models/voxcpm_refact/generate.rs b/src/models/voxcpm_refact/generate.rs new file mode 100644 index 0000000..1df5760 --- /dev/null +++ b/src/models/voxcpm_refact/generate.rs @@ -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>, + out_sample_rate: usize, + // model_name: String, +} + +impl VoxCPMGenerateRefact { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + 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 { + 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::()? 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>> { + 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::()? 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")), + } + } +} diff --git a/src/models/voxcpm_refact/mod.rs b/src/models/voxcpm_refact/mod.rs new file mode 100644 index 0000000..c47380f --- /dev/null +++ b/src/models/voxcpm_refact/mod.rs @@ -0,0 +1,3 @@ +pub mod generate; +pub mod model; +pub mod processor; diff --git a/src/models/voxcpm_refact/model.rs b/src/models/voxcpm_refact/model.rs new file mode 100644 index 0000000..9358a43 --- /dev/null +++ b/src/models/voxcpm_refact/model.rs @@ -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, + stop_proj: Linear, + stop_head: Linear, + device: Device, + dtype: DType, +} + +impl VoxCPMModelRefact { + pub fn new(vb: VarBuilder, config: VoxCPMConfig, latent_dim: usize) -> Result { + 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 { + 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::()? == 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::()?; + 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, + audio_mask: Option, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + audio_vae: &AudioVAE, + ) -> Result>> { + 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::()? == 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::()?; + 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(); + } +} diff --git a/src/models/voxcpm_refact/processor.rs b/src/models/voxcpm_refact/processor.rs new file mode 100644 index 0000000..00bdf60 --- /dev/null +++ b/src/models/voxcpm_refact/processor.rs @@ -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> { + 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, + prompt_wav_path: Option, + tokenizer: &SingleChineseTokenizer, + audio_vae: &AudioVAE, + ) -> Result<(Tensor, Option, Option)> { + 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, + tokenizer: &SingleChineseTokenizer, + ) -> Result<(Tensor, Option, Option)> { + 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)) + } +} diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index d399849..abd1131 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -366,7 +366,15 @@ pub fn get_audio_bytes_vec(path_str: &str) -> Result> { 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 + )), + } } } diff --git a/tests/test_voxcpm1_5.rs b/tests/test_voxcpm1_5.rs index 569b98a..5b5fd33 100644 --- a/tests/test_voxcpm1_5.rs +++ b/tests/test_voxcpm1_5.rs @@ -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(()) +}