From f374b77acd53912a00299ace45fe213b51ffa81e Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Sat, 25 Apr 2026 15:55:23 +0800 Subject: [PATCH] update voxcpm stream --- src/models/voxcpm_refact/model.rs | 39 ++++++++++++++++++++++++------- 1 file changed, 31 insertions(+), 8 deletions(-) diff --git a/src/models/voxcpm_refact/model.rs b/src/models/voxcpm_refact/model.rs index ad1906f..20de981 100644 --- a/src/models/voxcpm_refact/model.rs +++ b/src/models/voxcpm_refact/model.rs @@ -349,10 +349,19 @@ impl VoxCPMModelRefact { )?; (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 streaming_prefix_len = 4usize; + let mut pred_feat_seq = Vec::with_capacity(streaming_prefix_len); + if let Some(audio_feat) = &audio_feat + && let Some(audio_mask) = &audio_mask + && audio_mask.i((1, t - 1))?.to_scalar::()? == 1 + { + let audio_len = audio_mask.sum_all()?.to_scalar::()? 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; let enc_outputs = self @@ -422,10 +431,15 @@ impl VoxCPMModelRefact { .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 mut decode_audio = audio_vae - .decode(&single_feat_pred.to_dtype(DType::F32)?, None)? - .squeeze(1)?; + // 保持容量不超过最大值 + if pred_feat_seq.len() == streaming_prefix_len { + 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 @@ -434,6 +448,15 @@ impl VoxCPMModelRefact { .argmax(D::Minus1)? .i(0)? .to_scalar::()?; + + let pred_feat_chunk = Tensor::cat(&pred_feat_seq, 1)?; + let (b, _, _, d) = pred_feat_chunk.dims4()?; + let feat_pred = pred_feat_chunk.permute((0, 3, 1, 2))? + .reshape((b, d, ()))? + .contiguous()?; + let mut decode_audio = audio_vae + .decode(&feat_pred.to_dtype(DType::F32)?, None)? + .squeeze(1)?; if i > min_len && stop_flag == 1 { // 最后一段去除噪音 let decode_audio_len = decode_audio.dim(D::Minus1)? - 640;