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>(
|
pub fn generate_stream_generic<M: InferenceModel>(
|
||||||
model: &mut M,
|
model: &mut M,
|
||||||
tokenizer: &TokenizerModel,
|
tokenizer: &TokenizerModel,
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ use crate::{
|
|||||||
MultiModalData,
|
MultiModalData,
|
||||||
generate::{
|
generate::{
|
||||||
GenerationContext, generate_generic, generate_generic_text, generate_stream_generic,
|
GenerationContext, generate_generic, generate_generic_text, generate_stream_generic,
|
||||||
|
generate_stream_generic_text,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse},
|
||||||
@@ -206,6 +207,50 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
|||||||
&mut ctx,
|
&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> {
|
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||||
|
|||||||
@@ -425,7 +425,6 @@ impl VoxCPMModelRefact {
|
|||||||
let decode_audio = audio_vae
|
let decode_audio = audio_vae
|
||||||
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
|
.decode(&single_feat_pred.to_dtype(DType::F32)?, None)?
|
||||||
.squeeze(1)?;
|
.squeeze(1)?;
|
||||||
yield Ok(decode_audio);
|
|
||||||
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
|
||||||
@@ -435,8 +434,12 @@ impl VoxCPMModelRefact {
|
|||||||
.i(0)?
|
.i(0)?
|
||||||
.to_scalar::<u32>()?;
|
.to_scalar::<u32>()?;
|
||||||
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 = decode_audio.narrow(D::Minus1, 0, decode_audio_len)?;
|
||||||
|
yield Ok(decode_audio); // 最后一段去除噪音
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
yield Ok(decode_audio);
|
||||||
position_id += seq_len;
|
position_id += seq_len;
|
||||||
seq_len = 1;
|
seq_len = 1;
|
||||||
lm_hidden = self
|
lm_hidden = self
|
||||||
|
|||||||
Reference in New Issue
Block a user