Merge branch 'jhqxxx:main' into main

This commit is contained in:
PoRi
2026-05-08 11:39:45 +08:00
committed by GitHub
2 changed files with 40 additions and 19 deletions
+3 -1
View File
@@ -76,6 +76,7 @@ impl VoxCPMGenerateRefact {
// Some(128), // Some(128),
// Some(false), // Some(false),
)?; )?;
let decode_chunk_size = audio_config.decoder_rates.iter().product();
let processor = VoxCPMProcessor::new( let processor = VoxCPMProcessor::new(
audio_vae.sample_rate, audio_vae.sample_rate,
audio_vae.chunk_size, audio_vae.chunk_size,
@@ -105,7 +106,8 @@ impl VoxCPMGenerateRefact {
VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device) VarBuilder::from_tensors(dict_to_hashmap, m_dtype, device)
}; };
let tokenizer = SingleChineseTokenizer::new(path)?; 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 let out_sample_rate = audio_config
.out_sample_rate .out_sample_rate
.unwrap_or(audio_config.sample_rate); .unwrap_or(audio_config.sample_rate);
+36 -17
View File
@@ -17,6 +17,7 @@ pub struct VoxCPMModelRefact {
config: VoxCPMConfig, config: VoxCPMConfig,
patch_size: usize, patch_size: usize,
latent_dim: usize, latent_dim: usize,
decode_patch_len: usize,
// audio_start_token: u32, // audio_start_token: u32,
// // audio_end_token: u32, // // audio_end_token: u32,
// ref_audio_start_token: u32, // ref_audio_start_token: u32,
@@ -39,7 +40,12 @@ pub struct VoxCPMModelRefact {
} }
impl 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 base_lm = MiniCPMModel::new(vb.pp("base_lm"), config.lm_config.clone())?;
let mut residual_lm_config = 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.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 stop_head = linear_no_bias(config.lm_config.hidden_size, 2, vb.pp("stop_head"))?;
let patch_size = config.patch_size; let patch_size = config.patch_size;
let decode_patch_len = patch_size * decode_chunk_size;
Ok(Self { Ok(Self {
config, config,
patch_size, patch_size,
latent_dim, latent_dim,
decode_patch_len,
// audio_start_token: 101, // audio_start_token: 101,
// // audio_end_token: 102, // // audio_end_token: 102,
// ref_audio_start_token: 103, // ref_audio_start_token: 103,
@@ -351,16 +359,28 @@ impl VoxCPMModelRefact {
}; };
let streaming_prefix_len = 4usize; let streaming_prefix_len = 4usize;
// 流式处理固定4个结果使得VAE decode结果正常
let mut pred_feat_seq = Vec::with_capacity(streaming_prefix_len); let mut pred_feat_seq = Vec::with_capacity(streaming_prefix_len);
// 流式处理添加prompt块
if let Some(audio_feat) = &audio_feat if let Some(audio_feat) = &audio_feat
&& let Some(audio_mask) = &audio_mask && 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 audio_len = audio_mask.dim(1)?;
let context_len = audio_len.min(streaming_prefix_len - 1); if audio_mask
let start = audio_feat.dim(1)? - context_len; .narrow(1, audio_len - 1, 1)?
let last_feat = audio_feat.narrow(1, start, context_len)?; .squeeze(0)?
pred_feat_seq.push(last_feat); .squeeze(0)?
.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)?
.to_dtype(self.dtype)?;
pred_feat_seq.push(last_feat);
}
} }
let mut position_id = 0; let mut position_id = 0;
let mut seq_len = t; let mut seq_len = t;
@@ -436,10 +456,6 @@ impl VoxCPMModelRefact {
pred_feat_seq.remove(0); pred_feat_seq.remove(0);
} }
pred_feat_seq.push(pred_feat.unsqueeze(1)?); 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; prefix_feat_cond = pred_feat;
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?; let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
let stop_flag = self let stop_flag = self
@@ -457,17 +473,20 @@ impl VoxCPMModelRefact {
let mut decode_audio = audio_vae let mut decode_audio = audio_vae
.decode(&feat_pred.to_dtype(DType::F32)?, None)? .decode(&feat_pred.to_dtype(DType::F32)?, None)?
.squeeze(1)?; .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 { if i > min_len && stop_flag == 1 {
// 最后一段去除噪音 // // 最后一段容易有噪音,暂时不返回
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640; // let decode_audio_len = decode_audio.dim(D::Minus1)? - 640;
decode_audio = decode_audio.narrow(D::Minus1, 0, decode_audio_len)?; // decode_audio = decode_audio.narrow(D::Minus1, 0, decode_audio_len)?;
yield Ok(decode_audio); // yield Ok(decode_audio);
break; break;
} }
if first_flag { if first_flag {
// 去除初始噪音 // 去除初始噪音
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640; let decode_audio_len = decode_audio.dim(D::Minus1)? - 1280;
decode_audio = decode_audio.narrow(D::Minus1, 640, decode_audio_len)?; decode_audio = decode_audio.narrow(D::Minus1, 1280, decode_audio_len)?;
first_flag = false; first_flag = false;
} }
yield Ok(decode_audio); yield Ok(decode_audio);