fix voxcpm stream end noise

This commit is contained in:
jhqxxx
2026-04-25 14:23:41 +08:00
parent 60a0ad69be
commit f390f249b7
3 changed files with 119 additions and 1 deletions
+70
View File
@@ -191,6 +191,76 @@ pub fn generate_generic<M: InferenceModel>(
))
}
pub fn generate_stream_generic_text<M: InferenceModel>(
model: &mut M,
tokenizer: &TokenizerModel,
input_ids: Tensor,
data: MultiModalData,
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
repeat_penalty: Option<f32>,
repeat_last_n: Option<usize>,
seed: u64,
max_tokens: u32,
device: &Device,
) -> Result<impl Stream<Item = Result<String, anyhow::Error>>> {
let mut ctx = GenerationContext::new(
temperature,
top_p,
top_k,
repeat_penalty,
repeat_last_n,
seed,
input_ids.dim(1)?,
max_tokens,
device.clone(),
);
let mut error_tokens = Vec::new();
let eos_ids = model.stop_token_ids();
let stream = stream! {
let mut input_ids = input_ids;
let mut generated = Vec::new();
for _ in 0..ctx.sample_len {
let logits = if ctx.seqlen_offset == 0 {
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
} else {
model.forward_step(&input_ids, ctx.seqlen_offset)
}?;
let next_token = sample_and_push(&mut ctx, &logits, &mut generated)?;
// 解码(处理的累积)
let decode_ids = if error_tokens.is_empty() {
vec![next_token]
} else {
let mut ids = error_tokens.clone();
ids.push(next_token);
ids
};
let decoded = tokenizer.token_decode(decode_ids)?;
if decoded.contains("") {
error_tokens.push(next_token);
if error_tokens.len() > 3 {
error_tokens.clear();
}
input_ids = ctx.prepare_for_next_token(next_token)?;
continue;
}
error_tokens.clear();
yield Ok(decoded);
if eos_ids.contains(&next_token) {
break;
}
input_ids = ctx.prepare_for_next_token(next_token)?;
}
model.clear_cache();
};
Ok(stream)
}
pub fn generate_stream_generic<M: InferenceModel>(
model: &mut M,
tokenizer: &TokenizerModel,
+45
View File
@@ -3,6 +3,7 @@ use crate::{
MultiModalData,
generate::{
GenerationContext, generate_generic, generate_generic_text, generate_stream_generic,
generate_stream_generic_text,
},
},
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
@@ -206,6 +207,50 @@ impl<'a> Qwen3_5GenerateModel<'a> {
&mut ctx,
)
}
pub fn generate_stream_text(
&mut self,
mes: ChatCompletionParameters,
) -> Result<impl Stream<Item = Result<String, anyhow::Error>>> {
let mes_render = self.chat_template.apply_chat_template(&mes)?;
let (mes_text, pixel_values, image_grid_thw, pixel_values_video, video_grid_thw) =
if let Some(processor) = &self.pre_processor {
let input = processor.process_info(&mes, &mes_render)?;
(
input.replace_text,
input.pixel_values,
input.image_grid_thw,
input.pixel_values_video,
input.video_grid_thw,
)
} else {
(mes_render, None, None, None, None)
};
let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?;
let sample_len = mes.max_tokens.unwrap_or(1024);
let data_vec = vec![
pixel_values,
image_grid_thw,
pixel_values_video,
video_grid_thw,
];
let data = MultiModalData::new(data_vec);
let seed = mes.seed.unwrap_or(34562) as u64;
generate_stream_generic_text(
&mut self.qwen3_5,
&self.tokenizer,
input_ids,
data,
mes.temperature,
mes.top_p,
None,
self.repeat_penalty.into(),
self.repeat_last_n.into(),
seed,
sample_len,
&self.device,
)
}
}
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
+4 -1
View File
@@ -425,7 +425,6 @@ impl VoxCPMModelRefact {
let decode_audio = audio_vae
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
.squeeze(1)?;
yield Ok(decode_audio);
prefix_feat_cond = pred_feat;
let stop_flag = self.stop_proj.forward(&lm_hidden)?.silu()?;
let stop_flag = self
@@ -435,8 +434,12 @@ impl VoxCPMModelRefact {
.i(0)?
.to_scalar::<u32>()?;
if i > min_len && stop_flag == 1 {
let decode_audio_len = decode_audio.dim(D::Minus1)? - 640;
let decode_audio = decode_audio.narrow(D::Minus1, 0, decode_audio_len)?;
yield Ok(decode_audio); // 最后一段去除噪音
break;
}
yield Ok(decode_audio);
position_id += seq_len;
seq_len = 1;
lm_hidden = self