update voxcpm
This commit is contained in:
+158
-30
@@ -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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Self> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<Tensor> {
|
||||
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<HashMap<String, Tensor>> {
|
||||
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<String, Tensor>,
|
||||
min_len: usize,
|
||||
max_len: usize,
|
||||
inference_timesteps: usize,
|
||||
cfg_value: f64,
|
||||
retry_badcase: bool,
|
||||
retry_badcase_ratio_threshold: f64,
|
||||
) -> Result<Tensor> {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user