update voxcpm stream
This commit is contained in:
@@ -76,6 +76,7 @@ impl VoxCPMGenerateRefact {
|
||||
// Some(128),
|
||||
// Some(false),
|
||||
)?;
|
||||
let decode_chunk_size = audio_config.decoder_rates.iter().product();
|
||||
let processor = VoxCPMProcessor::new(
|
||||
audio_vae.sample_rate,
|
||||
audio_vae.chunk_size,
|
||||
@@ -105,7 +106,8 @@ impl VoxCPMGenerateRefact {
|
||||
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 voxcpm =
|
||||
VoxCPMModelRefact::new(vb_voxcpm, config, audio_vae.latent_dim, decode_chunk_size)?;
|
||||
let out_sample_rate = audio_config
|
||||
.out_sample_rate
|
||||
.unwrap_or(audio_config.sample_rate);
|
||||
|
||||
@@ -17,6 +17,7 @@ pub struct VoxCPMModelRefact {
|
||||
config: VoxCPMConfig,
|
||||
patch_size: usize,
|
||||
latent_dim: usize,
|
||||
decode_patch_len: usize,
|
||||
// audio_start_token: u32,
|
||||
// // audio_end_token: u32,
|
||||
// ref_audio_start_token: u32,
|
||||
@@ -39,7 +40,12 @@ pub struct VoxCPMModelRefact {
|
||||
}
|
||||
|
||||
impl VoxCPMModelRefact {
|
||||
pub fn new(vb: VarBuilder, config: VoxCPMConfig, latent_dim: usize) -> Result<Self> {
|
||||
pub fn new(
|
||||
vb: VarBuilder,
|
||||
config: VoxCPMConfig,
|
||||
latent_dim: usize,
|
||||
decode_chunk_size: 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;
|
||||
@@ -116,10 +122,12 @@ impl VoxCPMModelRefact {
|
||||
let stop_head = linear_no_bias(config.lm_config.hidden_size, 2, vb.pp("stop_head"))?;
|
||||
|
||||
let patch_size = config.patch_size;
|
||||
let decode_patch_len = patch_size * decode_chunk_size;
|
||||
Ok(Self {
|
||||
config,
|
||||
patch_size,
|
||||
latent_dim,
|
||||
decode_patch_len,
|
||||
// audio_start_token: 101,
|
||||
// // audio_end_token: 102,
|
||||
// ref_audio_start_token: 103,
|
||||
@@ -351,16 +359,20 @@ impl VoxCPMModelRefact {
|
||||
};
|
||||
|
||||
let streaming_prefix_len = 4usize;
|
||||
// 流式处理固定4个结果使得VAE decode结果正常
|
||||
let mut pred_feat_seq = Vec::with_capacity(streaming_prefix_len);
|
||||
// 流式处理添加prompt块
|
||||
if let Some(audio_feat) = &audio_feat
|
||||
&& let Some(audio_mask) = &audio_mask
|
||||
&& audio_mask.i((1, t - 1))?.to_scalar::<u32>()? == 1
|
||||
{
|
||||
let audio_len = audio_mask.sum_all()?.to_scalar::<u32>()? as usize;
|
||||
let context_len = audio_len.min(streaming_prefix_len - 1);
|
||||
let start = audio_feat.dim(1)? - context_len;
|
||||
let last_feat = audio_feat.narrow(1, start, context_len)?;
|
||||
pred_feat_seq.push(last_feat);
|
||||
let audio_len = audio_mask.dim(1)?;
|
||||
if audio_mask.narrow(1, audio_len - 1, 1)?.to_scalar::<u32>()? == 1 {
|
||||
let audio_len = audio_mask.sum_all()?.to_scalar::<u32>()? as usize;
|
||||
let context_len = audio_len.min(streaming_prefix_len - 1);
|
||||
let start = audio_feat.dim(1)? - context_len;
|
||||
let last_feat = audio_feat.narrow(1, start, context_len)?;
|
||||
pred_feat_seq.push(last_feat);
|
||||
}
|
||||
}
|
||||
let mut position_id = 0;
|
||||
let mut seq_len = t;
|
||||
@@ -436,10 +448,6 @@ impl VoxCPMModelRefact {
|
||||
pred_feat_seq.remove(0);
|
||||
}
|
||||
pred_feat_seq.push(pred_feat.unsqueeze(1)?);
|
||||
// let single_feat_pred = pred_feat.permute((0, 2, 1))?.contiguous()?;
|
||||
// let mut decode_audio = audio_vae
|
||||
// .decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
|
||||
// .squeeze(1)?;
|
||||
prefix_feat_cond = pred_feat;
|
||||
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
|
||||
let stop_flag = self
|
||||
@@ -457,6 +465,9 @@ impl VoxCPMModelRefact {
|
||||
let mut decode_audio = audio_vae
|
||||
.decode(&feat_pred.to_dtype(DType::F32)?, None)?
|
||||
.squeeze(1)?;
|
||||
// 只取当前帧结果
|
||||
let decode_start = decode_audio.dim(D::Minus1)? - self.decode_patch_len;
|
||||
decode_audio = decode_audio.narrow(D::Minus1, decode_start, self.decode_patch_len)?;
|
||||
if i > min_len && stop_flag == 1 {
|
||||
// 最后一段去除噪音
|
||||
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640;
|
||||
|
||||
Reference in New Issue
Block a user