diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs index 7c0ef1b..cd3502d 100644 --- a/src/models/common/generate.rs +++ b/src/models/common/generate.rs @@ -191,6 +191,76 @@ pub fn generate_generic( )) } +pub fn generate_stream_generic_text( + model: &mut M, + tokenizer: &TokenizerModel, + input_ids: Tensor, + data: MultiModalData, + temperature: Option, + top_p: Option, + top_k: Option, + repeat_penalty: Option, + repeat_last_n: Option, + seed: u64, + max_tokens: u32, + device: &Device, +) -> Result>> { + 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( model: &mut M, tokenizer: &TokenizerModel, diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs index 03e6d16..7b56354 100644 --- a/src/models/qwen3_5/generate.rs +++ b/src/models/qwen3_5/generate.rs @@ -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>> { + 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> { diff --git a/src/models/voxcpm_refact/model.rs b/src/models/voxcpm_refact/model.rs index 9358a43..6450eb4 100644 --- a/src/models/voxcpm_refact/model.rs +++ b/src/models/voxcpm_refact/model.rs @@ -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::()?; 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