update voxcpm stream

This commit is contained in:
jhqxxx
2026-04-25 15:55:23 +08:00
parent 3ce5d8d024
commit f374b77acd
+31 -8
View File
@@ -349,10 +349,19 @@ impl VoxCPMModelRefact {
)?; )?;
(text_embed, prefix_feat_cond, None) (text_embed, prefix_feat_cond, None)
}; };
// let mut pred_feat_seq = Vec::new();
// if feat_mask.i((1, t-1))?.to_scalar::<f32>()? == 0.0 { let streaming_prefix_len = 4usize;
// // TODO for stream 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::<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 position_id = 0;
let mut seq_len = t; let mut seq_len = t;
let enc_outputs = self let enc_outputs = self
@@ -422,10 +431,15 @@ impl VoxCPMModelRefact {
.transpose(1, 2)?; // [b, p, d] .transpose(1, 2)?; // [b, p, d]
let curr_embed = self.feat_encoder.forward(&pred_feat.unsqueeze(1)?)?; // [b, 1, c] 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 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 if pred_feat_seq.len() == streaming_prefix_len {
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)? pred_feat_seq.remove(0);
.squeeze(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
@@ -434,6 +448,15 @@ impl VoxCPMModelRefact {
.argmax(D::Minus1)? .argmax(D::Minus1)?
.i(0)? .i(0)?
.to_scalar::<u32>()?; .to_scalar::<u32>()?;
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 { 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;