fix voxcpm stream end noise
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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> {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user