diff --git a/src/models/base_modules/mod.rs b/src/models/base_modules/mod.rs index 774719e..77fd179 100644 --- a/src/models/base_modules/mod.rs +++ b/src/models/base_modules/mod.rs @@ -179,7 +179,7 @@ impl AttentionNobias { Ok(attn_output) } - pub fn forward_step( + pub fn forward_with_cache( &mut self, xs: &Tensor, cos: &Tensor, diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 6c5f644..4be620e 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -66,7 +66,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { None => 2048, }; for _ in 0..sample_len { - let logits = self.minicpm.forward_step(&input_ids, seqlen_offset)?; + let logits = self.minicpm.forward_with_cache(&input_ids, seqlen_offset)?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; generate.push(next_token); @@ -98,7 +98,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { let stream = stream! { let mut error_tokens = Vec::new(); for _ in 0..sample_len { - let logits = self.minicpm.forward_step( + let logits = self.minicpm.forward_with_cache( &input_ids, seqlen_offset, )?; diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index 7ffee13..8fe530a 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -159,7 +159,7 @@ impl MiniCPMDecoderLayer { Ok(xs) } - pub fn forward_step( + pub fn forward_with_cache( &mut self, xs: &Tensor, cos: &Tensor, @@ -168,7 +168,7 @@ impl MiniCPMDecoderLayer { ) -> Result { let residual = xs.clone(); let xs = self.input_layernorm.forward(xs)?; - let xs = self.self_attn.forward_step(&xs, cos, sin, attention_mask, true)?; + let xs = self.self_attn.forward_with_cache(&xs, cos, sin, attention_mask, true)?; let xs = (residual + xs.affine( self.scale_depth as f64 / (self.num_hidden_layers as f64).sqrt(), @@ -254,7 +254,7 @@ impl MiniCPMModel { Ok(logits) } - pub fn forward_step(&mut self, input_ids: &Tensor, position_id: usize) -> Result { + pub fn forward_with_cache(&mut self, input_ids: &Tensor, position_id: usize) -> Result { let (bs, seq_len) = input_ids.dims2()?; let input_embeds = self .embed_tokens @@ -276,7 +276,7 @@ impl MiniCPMModel { let mut hidden_states = input_embeds; for decode_layer in &mut self.layers { hidden_states = - decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?; + decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; } hidden_states = self.norm.forward(&hidden_states)?; let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?; diff --git a/src/models/voxcpm/audio_vae.rs b/src/models/voxcpm/audio_vae.rs index 0c478a4..b7b7395 100644 --- a/src/models/voxcpm/audio_vae.rs +++ b/src/models/voxcpm/audio_vae.rs @@ -1,7 +1,7 @@ use anyhow::{Ok, Result}; use candle_core::{D, Tensor}; use candle_nn::{Conv1d, Conv1dConfig, ConvTranspose1d, ConvTranspose1dConfig, Module, VarBuilder}; -use std::{result::Result::Ok as StdOk, thread, time}; +use std::{result::Result::Ok as StdOk}; pub struct CausalConv1d { conv1d: Conv1d, @@ -40,7 +40,6 @@ pub struct CausalConvTranspose1d { conv_transpose1d: ConvTranspose1d, padding: usize, output_padding: usize, - config: ConvTranspose1dConfig, } impl CausalConvTranspose1d { @@ -66,7 +65,6 @@ impl CausalConvTranspose1d { conv_transpose1d, padding, output_padding, - config }) } pub fn forward(&self, x: &Tensor) -> Result { @@ -468,10 +466,10 @@ impl CausalDecoder { } pub struct AudioVAE { - encoder_dim: usize, - encoder_rates: Vec, - decoder_dim: usize, - decoder_rates: Vec, + // encoder_dim: usize, + // encoder_rates: Vec, + // decoder_dim: usize, + // decoder_rates: Vec, pub latent_dim: usize, hop_length: usize, encoder: CausalEncoder, @@ -511,10 +509,10 @@ impl AudioVAE { )?; let chunk_size = hop_length; Ok(Self { - encoder_dim, - encoder_rates, - decoder_dim, - decoder_rates, + // encoder_dim, + // encoder_rates, + // decoder_dim, + // decoder_rates, latent_dim, hop_length, encoder, diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs new file mode 100644 index 0000000..6988f5d --- /dev/null +++ b/src/models/voxcpm/generate.rs @@ -0,0 +1,161 @@ +use std::collections::HashMap; + +use crate::{ + models::voxcpm::{ + audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel, + tokenizer::SingleChineseTokenizer, + }, + utils::utils::{find_type_files, get_device, get_dtype}, +}; +use anyhow::{Ok, Result}; +use candle_core::{DType, Device, Tensor, pickle::read_all_with_key}; +use candle_nn::VarBuilder; + +pub struct VoxCPMGenerate { + voxcpm: VoxCPMModel, + prompt_cache: Option>, +} + +impl VoxCPMGenerate { + 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 cfg_dtype = config.dtype.as_str(); + let model_list = find_type_files(path, "pth")?; + println!(" pth model_list: {:?}", model_list); + 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 { + // println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb_vae = VarBuilder::from_tensors(dict_to_hashmap, vae_dtype, &device); + let audio_vae = AudioVAE::new( + vb_vae, + 128, + vec![2, 5, 8, 8], + Some(64), + 1536, + vec![8, 8, 5, 2], + 16000, + )?; + + let model_list = find_type_files(path, "bin")?; + println!(" bin model_list: {:?}", model_list); + dict_to_hashmap = HashMap::new(); + let mut m_dtype = get_dtype(dtype, cfg_dtype); + for m in model_list { + let dict = read_all_with_key(m, Some("state_dict"))?; + m_dtype = dict[0].1.dtype(); + for (k, v) in dict { + // println!("key: {}, tensor shape: {:?}", k, v); + dict_to_hashmap.insert(k, v); + } + } + let vb_voxcpm = VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device); + let config_path = path.to_string() + "/config.json"; + let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + let tokenizer = SingleChineseTokenizer::new(path)?; + let voxcpm = VoxCPMModel::new(vb_voxcpm, config, tokenizer, audio_vae)?; + + Ok(Self { + voxcpm, + prompt_cache: None, + }) + } + + pub fn build_prompt_cache( + &mut self, + prompt_text: String, + prompt_wav_path: String, + ) -> Result<()> { + let cache = self + .voxcpm + .build_prompt_cache(prompt_text, prompt_wav_path)?; + self.prompt_cache = Some(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 prompt_cache = cache.clone(); + self.voxcpm.generate_with_prompt_cache( + target_text, + prompt_cache, + min_len, + max_len, + inference_timesteps, + cfg_value, + retry_badcase, + retry_badcase_ratio_threshold, + )? + } + None => self.generate_simple(target_text)?, + }; + Ok(audio) + } + + pub fn generate_with_prompt_simple( + &mut self, + target_text: String, + prompt_text: Option, + prompt_wav_path: Option, + ) -> Result { + let audio = self.generate( + target_text, + prompt_text, + prompt_wav_path, + 2, + 1000, + 10, + 2.0, + false, + 6.0, + )?; + Ok(audio) + } + pub fn generate_simple(&mut self, target_text: String) -> Result { + let audio = self.generate(target_text, None, None, 2, 1000, 10, 2.0, false, 6.0)?; + Ok(audio) + } + pub fn generate( + &mut self, + target_text: String, + prompt_text: Option, + prompt_wav_path: Option, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + retry_badcase: bool, + retry_badcase_ratio_threshold: f64, + ) -> Result { + let audio = self.voxcpm.generate( + target_text, + prompt_text, + prompt_wav_path, + min_len, + max_len, + inference_timesteps, + cfg_value, + retry_badcase, + retry_badcase_ratio_threshold, + )?; + Ok(audio) + } +} diff --git a/src/models/voxcpm/minicpm4.rs b/src/models/voxcpm/minicpm4.rs index b454010..e0e0e36 100644 --- a/src/models/voxcpm/minicpm4.rs +++ b/src/models/voxcpm/minicpm4.rs @@ -1,4 +1,3 @@ -use std::{thread, time}; use crate::{ models::{ @@ -10,7 +9,7 @@ use crate::{ }; use anyhow::{anyhow, Ok, Result}; use candle_core::{DType, Device, Tensor, D}; -use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_norm}; +use candle_nn::{Embedding, Module, RmsNorm, VarBuilder, embedding, rms_norm}; pub struct MiniCPMLongRoPE { short_factor: Vec, @@ -177,7 +176,7 @@ impl MiniCPMDecoderLayer { Ok(xs) } - pub fn forward_step( + pub fn forward_with_cache( &mut self, xs: &Tensor, cos: &Tensor, @@ -188,7 +187,7 @@ impl MiniCPMDecoderLayer { let xs = self.input_layernorm.forward(xs)?; let xs = self .self_attn - .forward_step(&xs, cos, sin, attention_mask, true)?; + .forward_with_cache(&xs, cos, sin, attention_mask, true)?; let xs = if self.use_mup { let res_add = (residual + xs.affine( @@ -221,12 +220,11 @@ impl MiniCPMDecoderLayer { } pub struct MiniCPMModel { - cfg: VoxMiniCPM4Config, + // cfg: VoxMiniCPM4Config, pub embed_tokens: Option, layers: Vec, norm: RmsNorm, rope_emb: MiniCPMLongRoPE, - // lm_head: Linear, } impl MiniCPMModel { @@ -250,23 +248,17 @@ impl MiniCPMModel { } let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?; let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device(), vb.dtype())?; - // let lm_head = Linear::new(embed_tokens.embeddings().clone(), None); Ok(Self { - cfg, + // cfg, embed_tokens, layers, norm, rope_emb, - // lm_head, }) } pub fn forward(&mut self, input_embeds: &Tensor, position_id: usize, is_causal: bool) -> Result { let (bs, seq_len, _) = input_embeds.dims3()?; - // let input_embeds = self - // .embed_tokens - // .forward(&input_ids)? - // .affine(self.cfg.scale_emb, 0.0)?; let attention_mask: Option<&Tensor> = { if !is_causal || seq_len <= 1 { None @@ -288,17 +280,13 @@ impl MiniCPMModel { Ok(hidden_states) } - pub fn forward_step(&mut self, input_embeds: &Tensor, position_id: usize) -> Result { + pub fn forward_with_cache(&mut self, input_embeds: &Tensor, position_id: usize) -> Result { let input_embeds = match input_embeds.rank() { 2 => input_embeds.unsqueeze(1)?, 3 => input_embeds.clone(), _ => return Err(anyhow!("MiniCPMModelinput_embeds illigal")) }; let (bs, seq_len, _) = input_embeds.dims3()?; - // let input_embeds = self - // .embed_tokens - // .forward(&input_ids)? - // .affine(self.cfg.scale_emb, 0.0)?; let attention_mask: Option<&Tensor> = { if seq_len <= 1 { None @@ -315,7 +303,7 @@ impl MiniCPMModel { let mut hidden_states = input_embeds.clone(); for decode_layer in &mut self.layers { hidden_states = - decode_layer.forward_step(&hidden_states, &cos, &sin, attention_mask)?; + decode_layer.forward_with_cache(&hidden_states, &cos, &sin, attention_mask)?; } hidden_states = self.norm.forward(&hidden_states)?; diff --git a/src/models/voxcpm/mod.rs b/src/models/voxcpm/mod.rs index 1695295..d8272e6 100644 --- a/src/models/voxcpm/mod.rs +++ b/src/models/voxcpm/mod.rs @@ -2,4 +2,5 @@ pub mod config; pub mod audio_vae; pub mod minicpm4; pub mod tokenizer; -pub mod model; \ No newline at end of file +pub mod model; +pub mod generate; \ No newline at end of file diff --git a/src/models/voxcpm/model.rs b/src/models/voxcpm/model.rs index cfd695e..3c2f145 100644 --- a/src/models/voxcpm/model.rs +++ b/src/models/voxcpm/model.rs @@ -1,4 +1,4 @@ -use std::{cmp::max, f64, thread, time}; +use std::{cmp::max, collections::HashMap, f64}; use anyhow::{Ok, Result}; use candle_core::{D, DType, Device, IndexOp, Tensor}; @@ -7,7 +7,7 @@ use candle_transformers::models::deepseek2::SplitOp; use crate::{ models::voxcpm::{ - audio_vae::{self, AudioVAE}, + audio_vae::{AudioVAE}, config::{CfmConfig, VoxCPMConfig, VoxMiniCPM4Config}, minicpm4::MiniCPMModel, tokenizer::SingleChineseTokenizer, @@ -118,8 +118,8 @@ pub struct VoxCPMLocDiT { time_mlp: TimestepEmbedding, delta_time_mlp: TimestepEmbedding, decoder: MiniCPMModel, - config: VoxMiniCPM4Config, - in_channels: usize, + // config: VoxMiniCPM4Config, + // in_channels: usize, } impl VoxCPMLocDiT { @@ -150,8 +150,8 @@ impl VoxCPMLocDiT { time_mlp, delta_time_mlp, decoder, - config, - in_channels, + // config, + // in_channels, }) } @@ -188,9 +188,9 @@ impl VoxCPMLocDiT { } pub struct UnifiedCFM { - solver: String, - sigma_min: f32, - t_scheduler: String, + // solver: String, + // sigma_min: f32, + // t_scheduler: String, in_channels: usize, mean_mode: bool, estimator: VoxCPMLocDiT, @@ -207,9 +207,9 @@ impl UnifiedCFM { let sigma_min = cfm_params.sigma_min; let t_scheduler = cfm_params.t_scheduler; Ok(Self { - solver, - sigma_min, - t_scheduler, + // solver, + // sigma_min, + // t_scheduler, in_channels, mean_mode, estimator, @@ -227,7 +227,7 @@ impl UnifiedCFM { sway_sampling_coef: f64, use_cfg_zero_star: bool, ) -> Result { - let (b, c) = mu.dims2()?; + let (b, _) = mu.dims2()?; let t = patch_size; let dtype = mu.dtype(); let z = Tensor::randn(0.0f32, 1.0, (b, self.in_channels, t), mu.device())? @@ -345,7 +345,7 @@ impl VoxCPMLocEnc { } pub fn forward(&mut self, x: &Tensor) -> Result { - let (b, t, p, d) = x.dims4()?; + let (b, t, _, _) = x.dims4()?; let x = self.in_proj.forward(x)?; let special_tokens = self.special_token.expand((b, t, 1, self.hidden_size))?; let x = Tensor::cat(&[special_tokens, x], 2)?; @@ -362,7 +362,7 @@ pub struct VoxCPMModel { config: VoxCPMConfig, patch_size: usize, audio_start_token: usize, - audio_end_token: usize, + // audio_end_token: usize, chunk_size: usize, sample_rate: usize, tokenizer: SingleChineseTokenizer, @@ -390,7 +390,7 @@ impl VoxCPMModel { ) -> Result { let base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?; let audio_start_token = 101usize; - let audio_end_token = 102usize; + // let audio_end_token = 102usize; 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; @@ -456,7 +456,7 @@ impl VoxCPMModel { config, patch_size, audio_start_token, - audio_end_token, + // audio_end_token, chunk_size: audio_vae.chunk_size, sample_rate: audio_vae.sample_rate, tokenizer, @@ -486,7 +486,6 @@ impl VoxCPMModel { inference_timesteps: usize, cfg_value: f64, retry_badcase: bool, - retry_badcase_max_times: usize, retry_badcase_ratio_threshold: f64, ) -> Result { let (text_token, text_mask, audio_feat, audio_mask) = match prompt_wav_path { @@ -554,17 +553,41 @@ impl VoxCPMModel { (text_token, text_mask, audio_feat, audio_mask) } }; - - let text_token = text_token.unsqueeze(0)?; - let text_mask = text_mask.unsqueeze(0)?; - let audio_feat = audio_feat.unsqueeze(0)?.to_dtype(DType::BF16)?; - let audio_mask = audio_mask.unsqueeze(0)?; let target_text_length = self.tokenizer.encode(target_text)?.len(); let max_len = if retry_badcase { (target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize } else { max_len }; + let decode_audio = self._generate( + &text_token, + &text_mask, + &audio_feat, + &audio_mask, + min_len, + max_len, + inference_timesteps, + cfg_value, + )?; + Ok(decode_audio) + } + + fn _generate( + &mut self, + text_token: &Tensor, + text_mask: &Tensor, + audio_feat: &Tensor, + audio_mask: &Tensor, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + ) -> Result { + let text_token = text_token.unsqueeze(0)?; + let text_mask = text_mask.unsqueeze(0)?; + let audio_feat = audio_feat.unsqueeze(0)?.to_dtype(DType::BF16)?; + let audio_mask = audio_mask.unsqueeze(0)?; + let latent_pred = self.inference( &text_token, &text_mask, @@ -584,7 +607,7 @@ impl VoxCPMModel { Ok(decode_audio) } - pub fn inference( + fn inference( &mut self, text: &Tensor, text_mask: &Tensor, @@ -595,7 +618,7 @@ impl VoxCPMModel { inference_timesteps: usize, cfg_value: f64, ) -> Result { - let (b, t, p, d) = feat.dims4()?; + let (_, t, _, _) = feat.dims4()?; let feat_embed = self.feat_encoder.forward(feat)?; // [b, t, h_feat] let feat_embed = self.enc_to_lm_proj.forward(&feat_embed)?; let scale_emb = if self.config.lm_config.use_mup { @@ -619,7 +642,7 @@ impl VoxCPMModel { let mut pred_feat_seq = Vec::new(); let mut position_id = 0; let mut seq_len = t; - let enc_outputs = self.base_lm.forward_step(&combined_embed, position_id)?; + let enc_outputs = self.base_lm.forward_with_cache(&combined_embed, position_id)?; let enc_outputs = self .fsq_layer .forward(&enc_outputs)? @@ -630,7 +653,7 @@ impl VoxCPMModel { let input_embeds = enc_outputs.add(&feat_mask.unsqueeze(D::Minus1)?.broadcast_mul(&feat_embed)?)?; - let residual_enc_outputs = self.residual_lm.forward_step(&input_embeds, position_id)?; + 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 { @@ -671,16 +694,16 @@ impl VoxCPMModel { seq_len = 1; lm_hidden = self .base_lm - .forward_step(&curr_embed.i((.., 0, ..))?, position_id)? + .forward_with_cache(&curr_embed.i((.., 0, ..))?, position_id)? .squeeze(1)?; lm_hidden = self.fsq_layer.forward(&lm_hidden)?; residual_hidden = self .residual_lm - .forward_step(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)? + .forward_with_cache(&lm_hidden.add(&curr_embed.i((.., 0, ..))?)?, position_id)? .squeeze(1)?; } let pred_seq = Tensor::cat(&pred_feat_seq, 1)?; // (b, t, p, d) - let (b, t, p, d) = pred_seq.dims4()?; + let (b, _, _, d) = pred_seq.dims4()?; let feat_pred = pred_seq .permute((0, 3, 1, 2))? .reshape((b, d, ()))? @@ -689,4 +712,109 @@ impl VoxCPMModel { self.residual_lm.clear_kv_cache(); Ok(feat_pred) } + + pub fn build_prompt_cache( + &mut self, + prompt_text: String, + prompt_wav_path: String, + ) -> Result> { + let text_token = self.tokenizer.encode(prompt_text)?; + let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?; + let mut audio = + load_audio_with_resample(prompt_wav_path, self.device.clone(), 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 = self.audio_vae.encode(&audio, Some(self.sample_rate))?; + let audio_feat = audio_feat + .reshape((self.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 generate_with_prompt_cache( + &mut self, + target_text: String, + prompt_cache: HashMap, + min_len: usize, + max_len: usize, + inference_timesteps: usize, + cfg_value: f64, + retry_badcase: bool, + retry_badcase_ratio_threshold: f64, + ) -> Result { + let target_text_token = self.tokenizer.encode(target_text.clone())?; + let target_text_token = + Tensor::from_slice(&target_text_token, target_text_token.len(), &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 as u32], &self.device)?; + let 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())), + None => (0, None), + }; + let (text_token, text_mask, audio_feat, audio_mask) = if audio_length > 0 { + let audio_feat = audio_feat.unwrap(); + let audio_length = audio_feat.dim(0)?; + let text_pad_token = Tensor::zeros(audio_length, DType::U32, &self.device)?; + let text_token = Tensor::cat(&[text_token, text_pad_token], D::Minus1)?; + let audio_pad_feat = Tensor::zeros( + (text_length, self.patch_size, self.audio_vae.latent_dim), + audio_feat.dtype(), + &self.device, + )?; + let audio_feat = Tensor::cat(&[audio_pad_feat, audio_feat], 0)?; + let text_mask = Tensor::cat( + &[ + Tensor::ones(text_length, self.dtype, &self.device)?, + Tensor::zeros(audio_length, self.dtype, &self.device)?, + ], + D::Minus1, + )?; + let audio_mask = Tensor::cat( + &[ + Tensor::zeros(text_length, self.dtype, &self.device)?, + Tensor::ones(audio_length, self.dtype, &self.device)?, + ], + D::Minus1, + )?; + (text_token, text_mask, audio_feat, audio_mask) + } else { + let audio_feat = Tensor::zeros( + (text_length, self.patch_size, self.audio_vae.latent_dim), + DType::F32, + &self.device, + )?; + let text_mask = Tensor::ones(text_length, self.dtype, &self.device)?; + let audio_mask = Tensor::zeros(text_length, self.dtype, &self.device)?; + (text_token, text_mask, audio_feat, audio_mask) + }; + let target_text_length = self.tokenizer.encode(target_text)?.len(); + let max_len = if retry_badcase { + (target_text_length as f64 * retry_badcase_ratio_threshold + 10.0) as usize + } else { + max_len + }; + let decode_audio = self._generate( + &text_token, + &text_mask, + &audio_feat, + &audio_mask, + min_len, + max_len, + inference_timesteps, + cfg_value, + )?; + Ok(decode_audio) + } } diff --git a/src/utils/audio_utils.rs b/src/utils/audio_utils.rs index 8f16e26..ccdc06a 100644 --- a/src/utils/audio_utils.rs +++ b/src/utils/audio_utils.rs @@ -226,7 +226,7 @@ pub fn load_audio>(path: P, device: Device) -> Result<(Tensor, us let samples: Vec = match spec.sample_format { SampleFormat::Int => { // 将整数样本转换为浮点数 [-1.0, 1.0] - println!("spec.bits_per_sample: {}", spec.bits_per_sample); + // println!("spec.bits_per_sample: {}", spec.bits_per_sample); let samples = match spec.bits_per_sample { 8 => { reader diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index f715399..c7b9b1a 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -223,7 +223,7 @@ pub fn masked_scatter_dim0(original: &Tensor, replace: &Tensor, mask: &Tensor) - let mask = mask.squeeze(0)?; let slices = nonzero_slice(&mask)?; let mut sub_start = 0usize; - let mut sub_end = 0usize; + let mut sub_end; for (start, end) in slices { sub_end = sub_start + (end - start); let sub_replace = replace.i((sub_start..sub_end, ..))?; diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs index f8a4308..836d17a 100644 --- a/tests/test_minicpm4.rs +++ b/tests/test_minicpm4.rs @@ -21,7 +21,7 @@ fn minicpm_generate() -> Result<()> { "messages": [ { "role": "user", - "content": "贾宝玉和孙悟空有什么关系" + "content": "你好啊,你是谁" } ] } diff --git a/tests/test_voxcpm.rs b/tests/test_voxcpm.rs index 93bd724..71d1c57 100644 --- a/tests/test_voxcpm.rs +++ b/tests/test_voxcpm.rs @@ -1,9 +1,9 @@ use anyhow::{Ok, Result}; -use std::collections::HashMap; +use std::{collections::HashMap, time::Instant}; use aha::{ models::voxcpm::{ - audio_vae::AudioVAE, config::VoxCPMConfig, model::VoxCPMModel, + audio_vae::AudioVAE, config::VoxCPMConfig, generate::VoxCPMGenerate, model::VoxCPMModel, tokenizer::SingleChineseTokenizer, }, utils::{ @@ -18,64 +18,47 @@ use candle_nn::VarBuilder; fn voxcpm_generate() -> Result<()> { // RUST_BACKTRACE=1 cargo test -F cuda,flash-attn voxcpm_generate -- --nocapture let model_path = "/home/jhq/huggingface_model/openbmb/VoxCPM-0.5B/"; - let model_list = find_type_files(&model_path, "pth")?; - println!(" pth model_list: {:?}", model_list); - let dev = get_device(None); - let mut dict_to_hashmap = HashMap::new(); - let mut dtype = candle_core::DType::F32; - for m in model_list { - let dict = read_all_with_key(m, Some("state_dict"))?; - dtype = dict[0].1.dtype(); - for (k, v) in dict { - // println!("key: {}, tensor shape: {:?}", k, v); - // if k.contains("decoder.model.2.block.1") { - // println!("val: {}", v); - // } - dict_to_hashmap.insert(k, v); - } - } - let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev); - let audio_vae = AudioVAE::new( - vb, - 128, - vec![2, 5, 8, 8], - Some(64), - 1536, - vec![8, 8, 5, 2], - 16000, + + let i_start = Instant::now(); + let mut voxcpm_generate = VoxCPMGenerate::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.generate( + // "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), + // Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), + // Some("./assets/audio/voice_01.wav".to_string()), + // // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), + // // Some("./assets/audio/voice_05.wav".to_string()), + // 2, + // 100, + // 10, + // 2.0, + // false, + // 6.0, + // )?; + + // 创建prompt_cache + let _ = voxcpm_generate.build_prompt_cache( + "啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string(), + "./assets/audio/voice_01.wav".to_string(), )?; - println!("audio vae load down"); - let model_list = find_type_files(&model_path, "bin")?; - println!(" bin model_list: {:?}", model_list); - dict_to_hashmap = HashMap::new(); - for m in model_list { - let dict = read_all_with_key(m, Some("state_dict"))?; - dtype = dict[0].1.dtype(); - for (k, v) in dict { - // println!("key: {}, tensor shape: {:?}", k, v); - dict_to_hashmap.insert(k, v); - } - } - let vb_vox = VarBuilder::from_tensors(dict_to_hashmap, dtype, &dev); - let config_path = model_path.to_string() + "/config.json"; - let config: VoxCPMConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; - let tokenizer = SingleChineseTokenizer::new(model_path)?; - let mut voxcpm = VoxCPMModel::new(vb_vox, config, tokenizer, audio_vae)?; - let generate = voxcpm.generate( + // 使用prompt_cache生成语音 + let generate = voxcpm_generate.generate_use_prompt_cache( "太阳当空照,花儿对我笑,小鸟说早早早".to_string(), - Some("啥子小师叔,打狗还要看主人,你再要继续,我,就是你的对手".to_string()), - Some("./assets/audio/voice_01.wav".to_string()), - // Some("一定被灰太狼给吃了,我已经为他准备好了花圈了".to_string()), - // Some("./assets/audio/voice_05.wav".to_string()), 2, 100, 10, 2.0, false, - 3, 6.0, )?; - let _ = save_wav(&generate, "voxcpm_init.wav")?; + + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + let _ = save_wav(&generate, "voxcpm.wav")?; Ok(()) } diff --git a/voxcpm.wav b/voxcpm.wav new file mode 100644 index 0000000..6245715 Binary files /dev/null and b/voxcpm.wav differ diff --git a/voxcpm_init.wav b/voxcpm_init.wav deleted file mode 100644 index 3669a44..0000000 Binary files a/voxcpm_init.wav and /dev/null differ