diff --git a/README.md b/README.md index e402dcc..dc9599a 100644 --- a/README.md +++ b/README.md @@ -29,11 +29,11 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an | Category | Models | |----------|--------| -| **Text** | Qwen3, MiniCPM4,
LFM2, LFM2.5 | +| **Text** | Qwen3, MiniCPM4, LFM2, LFM2.5 | | **Vision** | Qwen2.5-VL, Qwen3-VL, Qwen3.5,
LFM2.5-VL, LFM2-VL | -| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 ,
PaddleOCR-VL, PaddleOCR-VL1.5,
Hunyuan-OCR, GLM-OCR | +| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 , PaddleOCR-VL
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR | | **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR | -| **Audio** | VoxCPM, VoxCPM1.5 | +| **TTS** | VoxCPM, VoxCPM1.5 | | **Image** | RMBG-2.0 (background removal) | ## Why aha? @@ -46,6 +46,11 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an - **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing ## Changelog +### 2026-04-02 +- refactor generate code +- \...\ The content of the thought chain is returned using the reasoning_content field. +- chat response add time info + ### 2026-04-01 - refactor deepseek_ocr/fun_asr_nano generate code @@ -64,9 +69,6 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an - add LFM2.5-1.2B-Instruct - add LFM2-1.2B -### v0.2.3 (2026-03-18) -- add DeepSeek-OCR-2 - **[View full changelog](docs/changelog.md)** → diff --git a/README.zh-CN.md b/README.zh-CN.md index a624d72..386e946 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -28,11 +28,11 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | 类别 | 模型 | |------|------| -| **文本** | Qwen3, MiniCPM4,
LFM2, LFM2.5 | +| **文本** | Qwen3, MiniCPM4, LFM2, LFM2.5 | | **视觉** | Qwen2.5-VL, Qwen3-VL, Qwen3.5
LFM2.5-VL, LFM2-VL | -| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2 ,
PaddleOCR-VL, PaddleOCR-VL1.5,
Hunyuan-OCR, GLM-OCR | +| **OCR** | DeepSeek-OCR, DeepSeek-OCR-2, PaddleOCR-VL,
PaddleOCR-VL1.5, Hunyuan-OCR, GLM-OCR | | **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR | -| **音频** | VoxCPM, VoxCPM1.5 | +| **TTS** | VoxCPM, VoxCPM1.5 | | **图像** | RMBG-2.0 (背景移除) | ## 为什么选择 aha? @@ -45,6 +45,12 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 - **🧠 注意力优化** - 可选 Flash Attention 支持,优化长序列处理 ## 更新日志 +## Changelog +### 2026-04-02 +- 重构生成代码 +- \...\ 思维链内容使用reasoning_content字段返回。 +- 对话返回添加耗时信息 + ### 2026-04-01 - 重构 deepseek_ocr/fun_asr_nano 生成代码 diff --git a/docs/changelog.md b/docs/changelog.md index 5a3bc11..bc980ed 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -5,6 +5,11 @@ All notable changes to aha will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +### 2026-04-02 +- refactor generate code +- \...\ The content of the thought chain is returned using the reasoning_content field. +- response add time info + ### 2026-04-01 - refactor deepseek_ocr/fun_asr_nano generate code diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index 8585629..4d5bf0c 100644 --- a/docs/changelog.zh-CN.md +++ b/docs/changelog.zh-CN.md @@ -5,6 +5,11 @@ 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/), 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 +### 2026-04-02 +- 重构生成代码 +- \...\ 思维链内容使用reasoning_content字段返回。 +- 对话返回添加耗时信息 + ### 2026-04-01 - 重构 deepseek_ocr/fun_asr_nano 生成代码 diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs index 5a3b269..c680d5f 100644 --- a/src/chat_template/mod.rs +++ b/src/chat_template/mod.rs @@ -133,6 +133,9 @@ impl<'a> ChatTemplate<'a> { pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result { let enable_thinking = extract_metadata_value::(&messages.metadata, "enable_thinking"); + let mes_thinking_param = messages.enable_thinking; + let enable_thinking = + Some(enable_thinking.unwrap_or(false) || mes_thinking_param.unwrap_or(false)); let context = context! { messages => &messages.messages, tools => &messages.tools.as_ref(), diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs index eb9c079..1f34d8b 100644 --- a/src/models/common/generate.rs +++ b/src/models/common/generate.rs @@ -9,7 +9,10 @@ use crate::{ models::common::{InferenceModel, MultiModalData}, params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse}, tokenizer::TokenizerModel, - utils::{build_completion_chunk_response, build_completion_response_with_time}, + utils::response_utils::{ + build_chunk_response_with_reasoning, build_chunk_response_with_usage, + build_completion_chunk_response, build_completion_response_with_time, + }, }; pub fn get_logit_processor( temperature: Option, @@ -98,6 +101,17 @@ fn sample_and_push( Ok(token) } +// TODO +// let logits = if self.repeat_penalty == 1. { +// logits +// } else { +// let start_at = generate.len().saturating_sub(self.repeat_last_n); +// candle_transformers::utils::apply_repeat_penalty( +// &logits, +// self.repeat_penalty, +// &generate[start_at..], +// )? +// }; pub fn generate_generic( model: &mut M, tokenizer: &TokenizerModel, @@ -155,6 +169,7 @@ pub fn generate_stream_generic( top_k: Option, seed: u64, max_tokens: u32, + in_reasoning: bool, device: &Device, model_name: &str, ) -> Result>> { @@ -167,14 +182,23 @@ pub fn generate_stream_generic( max_tokens, device.clone(), ); + let prompt_tokens = ctx.seq_len as u32; + let mut prompt_secs = 0.0f64; + let mut completion_tokens = 0u32; + let mut completion_secs = 0.0f64; let mut error_tokens = Vec::new(); let eos_ids = model.stop_token_ids(); let stream = stream! { let mut input_ids = input_ids; + let mut tool_call_id = None; + let mut tool_call_content = String::new(); + let mut in_reasoning = in_reasoning; // 处理 unicode 错误累积 for _ in 0..ctx.sample_len { + let i_start = Instant::now(); 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) }?; @@ -183,6 +207,13 @@ pub fn generate_stream_generic( let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; ctx.logit_processor.sample(&logits)? }; + completion_tokens += 1; + let i_duration = i_start.elapsed(); + if ctx.seqlen_offset == 0 { + prompt_secs += i_duration.as_secs_f64(); + } else { + completion_secs += i_duration.as_secs_f64(); + }; // 解码(处理�的累积) let decode_ids = if error_tokens.is_empty() { @@ -204,9 +235,60 @@ pub fn generate_stream_generic( continue; } error_tokens.clear(); - yield Ok(build_completion_chunk_response(decoded, model_name, None, None)); + if decoded.eq("") { + in_reasoning = true; + input_ids = ctx.prepare_for_next_token(next_token)?; + continue; + } + if decoded.eq("") { + in_reasoning = false; + input_ids = ctx.prepare_for_next_token(next_token)?; + continue; + } + // 处理特殊标记和工具调用 + match decoded.as_str() { + "" => { + // 开始工具调用 + tool_call_id = Some(uuid::Uuid::new_v4().to_string()); + input_ids = ctx.prepare_for_next_token(next_token)?; + continue; + } + "" => { + // 结束工具调用 + let chunk = build_completion_chunk_response( + decoded, + model_name, + tool_call_id.clone(), + Some(tool_call_content.clone()) + ); + tool_call_id = None; + tool_call_content = String::new(); + yield Ok(chunk); + } + _ => { + if tool_call_id.is_some() { + // 在工具调用过程中,收集工具调用内容 + tool_call_content.push_str(&decoded); + input_ids = ctx.prepare_for_next_token(next_token)?; + continue; + } else { + + // 正常文本输出 + let chunk = if in_reasoning { + build_chunk_response_with_reasoning(decoded, model_name) + } else { + build_completion_chunk_response( + decoded, model_name, + None, + None + )}; + yield Ok(chunk); + } + } + } if eos_ids.contains(&next_token) { + yield Ok(build_chunk_response_with_usage(model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into())); break; } input_ids = ctx.prepare_for_next_token(next_token)?; diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 7065ec7..19c19a0 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -18,14 +18,18 @@ impl MultiModalData { } } +#[allow(unused)] pub trait InferenceModel { /// 初始前向传播(考虑多模态输入) + /// 默认实现无特殊数据 fn forward_initial( &mut self, input_ids: &Tensor, seqlen_offset: usize, data: MultiModalData, - ) -> Result; + ) -> Result { + Self::forward_step(self, input_ids, seqlen_offset) + } /// 后续前向传播(自回归步骤) fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result; diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index 79eee51..939afaf 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -171,6 +171,7 @@ impl GenerateModel for DeepseekOCRGenerateModel { None, seed, max_tokens, + false, &self.device, &self.model_name, )?; diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index f97c04c..446648d 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -30,8 +30,6 @@ pub struct FunAsrNanoGenerateModel { fun_asr_nano: FunAsrNanoModel, device: Device, dtype: DType, - // eos_token_id1: u32, - // eos_token_id2: u32, generation_config: Qwen3GenerationConfig, model_name: String, } @@ -90,8 +88,6 @@ impl FunAsrNanoGenerateModel { fun_asr_nano, device, dtype, - // eos_token_id1: generation_config.eos_token_id[0] as u32, - // eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, model_name, }) @@ -165,58 +161,10 @@ impl GenerateModel for FunAsrNanoGenerateModel { top_k.into(), seed, max_tokens, + false, &self.device, &self.model_name, )?; - // let mut seq_len = input_ids.dim(1)?; - // let mut seqlen_offset = 0; - // let sample_len = mes.max_tokens.unwrap_or(1024); - // let stream = stream! { - // let mut error_tokens = Vec::new(); - // let mut speech = Some(speech.to_dtype(self.dtype)?); - // let mut fbank_mask = Some(&fbank_mask); - // let mut input_ids = input_ids; - // for _ in 0..sample_len { - // let logits = self.fun_asr_nano.forward( - // &input_ids, - // speech.as_ref(), - // fbank_mask, - // seqlen_offset, - // )?; - // let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - // let next_token = logit_processor.sample(&logits)?; - // let mut decode_ids = Vec::new(); - // if !error_tokens.is_empty() { - // decode_ids.extend_from_slice(&error_tokens); - // } - // decode_ids.push(next_token); - // let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - // if decoded_token.contains("�") { - // error_tokens.push(next_token); - // if error_tokens.len() > 3 { - // error_tokens.clear(); - // } - // seqlen_offset += seq_len; - // seq_len = 1; - // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - // speech = None; - // fbank_mask = None; - // continue; - // } - // error_tokens.clear(); - // let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - // yield Ok(chunk); - // if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - // break; - // } - // seqlen_offset += seq_len; - // seq_len = 1; - // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - // speech = None; - // fbank_mask = None; - // } - // self.fun_asr_nano.clear_kv_cache(); - // }; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/fun_asr_nano/model.rs b/src/models/fun_asr_nano/model.rs index 2d3118e..484e4b2 100644 --- a/src/models/fun_asr_nano/model.rs +++ b/src/models/fun_asr_nano/model.rs @@ -611,7 +611,7 @@ impl FunAsrNanoModel { config.audio_adaptor_conf.n_layer, 8, )?; - let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?; + let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"), vec![])?; Ok(Self { audio_encoder, audio_adaptor, diff --git a/src/models/fun_asr_nano/processor.rs b/src/models/fun_asr_nano/processor.rs index 662409b..04adcff 100644 --- a/src/models/fun_asr_nano/processor.rs +++ b/src/models/fun_asr_nano/processor.rs @@ -1,5 +1,5 @@ use crate::params::chat::ChatCompletionParameters; -use anyhow::Result; +use anyhow::{Result, anyhow}; use candle_core::{D, Device, Tensor}; use crate::{ @@ -92,6 +92,9 @@ impl FunAsrNanoProcessor { source_ids.extend_from_slice(&sub_token); fbank_mask.extend_from_slice(&vec![0u32; sub_token.len()]); let audio_tensors = extract_audios(mes, &self.device, Some(self.fronted_conf.fs))?; + if audio_tensors.is_empty() { + return Err(anyhow!("FunASRNano need audio input")); + } let audio = &audio_tensors[0]; let (speech, speech_lengths) = self.extract_fbank(audio)?; let olens = 1 + (speech_lengths - 3 + 2) / 2; diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index a28987e..66eada0 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -1,11 +1,13 @@ use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; -use anyhow::{Result, anyhow}; +use anyhow::Result; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::{ @@ -17,10 +19,7 @@ use crate::{ }, }, tokenizer::TokenizerModel, - utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, - }, + utils::{find_type_files, get_device, get_dtype}, }; pub struct GlmAsrNanoGenerateModel<'a> { @@ -30,9 +29,6 @@ pub struct GlmAsrNanoGenerateModel<'a> { glm_asr_nano: GlmAsrNanoModel, device: Device, dtype: DType, - eos_token_id1: u32, - eos_token_id2: u32, - eos_token_id3: u32, model_name: String, } @@ -48,7 +44,8 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { let dtype = get_dtype(dtype, cfg_dtype); let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?; + let eos_ids = vec![59246u32, 59253, 59255]; + let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg, eos_ids)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -61,9 +58,6 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { glm_asr_nano, device, dtype, - eos_token_id1: 59246, - eos_token_id2: 59253, - eos_token_id3: 59255, model_name, }) } @@ -72,46 +66,33 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let render_text: String = self.chat_template.apply_chat_template(&mes)?; let (input_features, audio_token_lengths, replace_text) = self.processor.process_info(&mes, &render_text)?; - let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - let mut input_features = Some(input_features.to_dtype(self.dtype)?); - let mut audio_token_lengths = Some(audio_token_lengths); - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate: Vec = Vec::new(); + let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; + let input_features = input_features.to_dtype(self.dtype)?; + let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logits = self.glm_asr_nano.forward( - input_features.as_ref(), - audio_token_lengths.as_ref(), - &input_ids, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.eos_token_id1 - || next_token == self.eos_token_id2 - || next_token == self.eos_token_id3 - { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - input_features = None; - audio_token_lengths = None; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.glm_asr_nano.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let mut ctx = GenerationContext::new( + mes.temperature, + mes.top_p, + mes.top_k, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + let data_vec = vec![input_features.into(), audio_token_lengths.into()]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.glm_asr_nano, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -126,58 +107,29 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> { >, > { let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let render_text = self.chat_template.apply_chat_template(&mes)?; let (input_features, audio_token_lengths, replace_text) = self.processor.process_info(&mes, &render_text)?; let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; + let input_features = input_features.to_dtype(self.dtype)?; + let audio_token_lengths = Tensor::new(audio_token_lengths, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut error_tokens = Vec::new(); - let mut input_features = Some(input_features.to_dtype(self.dtype)?); - let mut audio_token_lengths = Some(audio_token_lengths); - let mut input_ids = input_ids; - for _ in 0..sample_len { - let logits = - self.glm_asr_nano - .forward(input_features.as_ref(), audio_token_lengths.as_ref(), &input_ids, seqlen_offset)?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - input_features = None; - audio_token_lengths = None; - continue; - } - error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 || next_token == self.eos_token_id3{ - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - input_features = None; - audio_token_lengths = None; - } - self.glm_asr_nano.clear_kv_cache(); - }; + let data_vec = vec![input_features.into(), audio_token_lengths.into()]; + let data = MultiModalData::new(data_vec); + let stream = generate_stream_generic( + &mut self.glm_asr_nano, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs index 67d7285..e6d0a68 100644 --- a/src/models/glm_asr_nano/model.rs +++ b/src/models/glm_asr_nano/model.rs @@ -4,8 +4,11 @@ use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, linear_no use crate::{ models::{ - common::modules::{ - LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + common::{ + InferenceModel, + modules::{ + LlamaForCausalLM, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + }, }, glm_asr_nano::config::{GlmAsrAudioConfig, GlmAsrNanoConfig}, }, @@ -233,10 +236,11 @@ pub struct GlmAsrNanoModel { audio_tower: GlmAsrEncoder, multi_modal_projector: TwoLinearMLP, language_model: LlamaForCausalLM, + stop_token_ids: Vec, } impl GlmAsrNanoModel { - pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig) -> Result { + pub fn new(vb: VarBuilder, config: GlmAsrNanoConfig, eos_ids: Vec) -> Result { let audio_tower = GlmAsrEncoder::new(vb.pp("audio_tower"), &config.audio_config)?; let multi_modal_projector = TwoLinearMLP::new( vb.pp("multi_modal_projector"), @@ -273,6 +277,7 @@ impl GlmAsrNanoModel { audio_tower, multi_modal_projector, language_model, + stop_token_ids: eos_ids, }) } @@ -300,7 +305,7 @@ impl GlmAsrNanoModel { pub fn forward( &mut self, input_features: Option<&Tensor>, - audio_token_lengths: Option<&Vec>, + audio_token_lengths: Option<&Tensor>, input_ids: &Tensor, seqlen_offset: usize, ) -> Result { @@ -308,8 +313,9 @@ impl GlmAsrNanoModel { if let Some(input_features) = input_features && let Some(audio_token_len) = audio_token_lengths { + let audio_token_len = audio_token_len.to_vec1::()?; let audio_token_mask = get_equal_mask(input_ids, self.config.audio_token_id)?; - let audio_embeds = self.get_audio_features(input_features, audio_token_len)?; + let audio_embeds = self.get_audio_features(input_features, &audio_token_len)?; inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embeds, &audio_token_mask)?; } let logits = self.language_model.forward(&inputs_embeds, seqlen_offset)?; @@ -319,3 +325,38 @@ impl GlmAsrNanoModel { self.language_model.clear_kv_cache(); } } + +impl InferenceModel for GlmAsrNanoModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 2 { + return Err(anyhow::anyhow!( + "GlmAsrNano process data error, must have input_features, audio_token_lengths" + )); + } + let input_features = &data.data_vec[0]; + let audio_token_lengths = &data.data_vec[1]; + self.forward( + input_features.as_ref(), + audio_token_lengths.as_ref(), + input_ids, + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(None, None, input_ids, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/glm_asr_nano/processor.rs b/src/models/glm_asr_nano/processor.rs index 7b3f818..184e198 100644 --- a/src/models/glm_asr_nano/processor.rs +++ b/src/models/glm_asr_nano/processor.rs @@ -206,6 +206,9 @@ impl GlmAsrNanoProcessor { render_text: &str, ) -> Result<(Tensor, Vec, String)> { let audio_tensors = extract_audios(mes, &self.device, Some(self.sampling_rate))?; + if audio_tensors.is_empty() { + return Err(anyhow::anyhow!("GlmASRNano need audio input")); + } let (input_features, input_features_mask, per_sample_windows) = self.process_audio(audio_tensors)?; let audio_lengths = input_features_mask.sum(D::Minus1)?; diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs index caad0e2..e3b3c94 100644 --- a/src/models/glm_ocr/generate.rs +++ b/src/models/glm_ocr/generate.rs @@ -1,12 +1,14 @@ //! GLM-OCR Inference and Generation use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; use anyhow::{Result, anyhow}; -use candle_core::{DType, Device, IndexOp, Tensor}; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::{ @@ -21,8 +23,7 @@ use crate::{ }, tokenizer::TokenizerModel, utils::{ - build_completion_chunk_response, build_completion_response, extract_user_text, - find_type_files, get_device, get_dtype, img_utils::extract_image_url, + extract_user_text, find_type_files, get_device, get_dtype, img_utils::extract_image_url, }, }; @@ -32,7 +33,6 @@ pub struct GlmOcrGenerateModel { processor: GlmOcrProcessor, model: GlmOcrModel, device: Device, - eos_token_ids: Vec, model_name: String, image_token_id: u32, image_start_token_id: u32, @@ -54,10 +54,10 @@ impl GlmOcrGenerateModel { let processor = GlmOcrProcessor::new(path, &device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let model = GlmOcrModel::new(vb, cfg.clone())?; let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: GlmOcrGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let model = GlmOcrModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -69,7 +69,6 @@ impl GlmOcrGenerateModel { processor, model, device, - eos_token_ids: generation_config.eos_token_id.clone(), model_name, image_token_id: cfg.image_token_id, image_start_token_id: cfg.image_start_token_id, @@ -84,8 +83,6 @@ impl GlmOcrGenerateModel { impl GenerateModel for GlmOcrGenerateModel { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); - // Extract image path and prompt from messages let image_urls = extract_image_url(&mes); let image_path = image_urls @@ -110,55 +107,32 @@ impl GenerateModel for GlmOcrGenerateModel { self.spatial_merge_size, )?; - let mut input_ids = processed.input_ids; - let pixel_values = Some(processed.pixel_values); - let image_grid_thw = Some(processed.grid_thw); - let image_mask = Some(processed.image_mask); - let mut seqlen_offset = 0; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut generate = Vec::new(); - let sample_len = mes.max_tokens.unwrap_or(512); + let input_ids = processed.input_ids; + let sample_len = mes.max_tokens.unwrap_or(1024); + let mut ctx = GenerationContext::new( + mes.temperature, + mes.top_p, + mes.top_k, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); - for _ in 0..sample_len { - let is_first_pass = seqlen_offset == 0; - let logits = self.model.forward( - &input_ids, - if is_first_pass { - pixel_values.as_ref() - } else { - None - }, - if is_first_pass { - image_grid_thw.as_ref() - } else { - None - }, - if is_first_pass { - image_mask.as_ref() - } else { - None - }, - seqlen_offset, - )?; - let logits = logits.i((0, seq_len - 1, ..))?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - - generate.push(next_token); - if self.eos_token_ids.contains(&next_token) { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - } - - self.model.clear_kv_cache(); - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let data_vec = vec![ + processed.pixel_values.into(), + processed.grid_thw.into(), + processed.image_mask.into(), + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -173,8 +147,6 @@ impl GenerateModel for GlmOcrGenerateModel { >, > { let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); - // Extract image path and prompt from messages let image_urls = extract_image_url(&mes); let image_path = image_urls @@ -199,63 +171,28 @@ impl GenerateModel for GlmOcrGenerateModel { self.spatial_merge_size, )?; - let mut input_ids = processed.input_ids; - let pixel_values = Some(processed.pixel_values); - let image_grid_thw = Some(processed.grid_thw); - let image_mask = Some(processed.image_mask); - let mut seqlen_offset = 0; - let mut seq_len = input_ids.dim(1)?; - let sample_len = mes.max_tokens.unwrap_or(512); - - let stream = stream! { - let mut generated: Vec = Vec::new(); - let mut error_tokens = Vec::new(); - for _ in 0..sample_len { - let is_first_pass = seqlen_offset == 0; - let logits = self.model.forward( - &input_ids, - if is_first_pass { pixel_values.as_ref() } else { None }, - if is_first_pass { image_grid_thw.as_ref() } else { None }, - if is_first_pass { image_mask.as_ref() } else { None }, - seqlen_offset, - ).map_err(|e| anyhow!(format!("forward error: {e}")))?; - let logits = logits.i((0, seq_len - 1, ..)).map_err(|e| anyhow!(format!("index error: {e}")))?.to_dtype(DType::F32).map_err(|e| anyhow!(format!("dtype error: {e}")))?; - - let next_token = logit_processor.sample(&logits).map_err(|e| anyhow!(format!("sample error: {e}")))?; - generated.push(next_token); - - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("decode error: {e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device).map_err(|e| anyhow!(format!("tensor error: {e}")))?; - continue; - } - error_tokens.clear(); - - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - - if self.eos_token_ids.contains(&next_token) { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device).map_err(|e| anyhow!(format!("tensor error: {e}")))?; - } - self.model.clear_kv_cache(); - }; - + let input_ids = processed.input_ids; + let sample_len = mes.max_tokens.unwrap_or(1024); + let data_vec = vec![ + processed.pixel_values.into(), + processed.grid_thw.into(), + processed.image_mask.into(), + ]; + let data = MultiModalData::new(data_vec); + let stream = generate_stream_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/glm_ocr/model.rs b/src/models/glm_ocr/model.rs index 912084f..a650461 100644 --- a/src/models/glm_ocr/model.rs +++ b/src/models/glm_ocr/model.rs @@ -9,7 +9,7 @@ use candle_nn::{ use crate::{ models::{ - common::modules::GateUpDownMLP, + common::{InferenceModel, modules::GateUpDownMLP}, glm_ocr::config::{GlmOcrConfig, GlmOcrTextConfig, GlmOcrVisionConfig}, }, position_embed::rope::{apply_rotary_pos_emb_vision, glm_ocr_apply_rotary_pos_emb}, @@ -1256,7 +1256,8 @@ impl GlmOcrTextModel { } hidden_states = self.norm.forward(&hidden_states)?; - let logits = self.lm_head.forward(&hidden_states)?; + let last = hidden_states.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&last)?; Ok(logits) } @@ -1271,10 +1272,11 @@ impl GlmOcrTextModel { pub struct GlmOcrModel { vision_encoder: GlmOcrVisionModel, language_model: GlmOcrTextModel, + stop_token_ids: Vec, } impl GlmOcrModel { - pub fn new(vb: VarBuilder, config: GlmOcrConfig) -> Result { + pub fn new(vb: VarBuilder, config: GlmOcrConfig, eos_ids: Vec) -> Result { let vision_encoder = GlmOcrVisionModel::new(vb.pp("model").pp("visual"), &config.vision_config)?; let language_model = GlmOcrTextModel::new( @@ -1286,6 +1288,7 @@ impl GlmOcrModel { Ok(Self { vision_encoder, language_model, + stop_token_ids: eos_ids, }) } @@ -1331,3 +1334,40 @@ impl GlmOcrModel { self.language_model.clear_kv_cache(); } } + +impl InferenceModel for GlmOcrModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 3 { + return Err(anyhow::anyhow!( + "GlmOcr process data error, must have pixel_values, image_grid_thw, image_mask" + )); + } + let pixel_values = &data.data_vec[0]; + let image_grid_thw = &data.data_vec[1]; + let image_mask = &data.data_vec[2]; + self.forward( + input_ids, + pixel_values.as_ref(), + image_grid_thw.as_ref(), + image_mask.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/hunyuan_ocr/config.rs b/src/models/hunyuan_ocr/config.rs index 7cc172c..46bcbb9 100644 --- a/src/models/hunyuan_ocr/config.rs +++ b/src/models/hunyuan_ocr/config.rs @@ -85,7 +85,7 @@ pub struct HunyuanOCRGenerationConfig { pub bos_token_id: usize, pub pad_token_id: usize, pub do_sample: bool, - pub eos_token_id: Vec, + pub eos_token_id: Vec, pub top_p: f32, pub top_k: usize, pub temperature: f32, diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index 1a993d1..8b0e212 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -1,11 +1,13 @@ use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; -use anyhow::{Result, anyhow}; -use candle_core::{DType, Device, Tensor}; +use anyhow::Result; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::{ @@ -19,10 +21,7 @@ use crate::{ }, }, tokenizer::TokenizerModel, - utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, - }, + utils::{find_type_files, get_device, get_dtype}, }; pub struct HunyuanOCRGenerateModel<'a> { @@ -31,8 +30,6 @@ pub struct HunyuanOCRGenerateModel<'a> { pre_processor: HunyuanVLProcessor, hunyuan_vl: HunyuanVLModel, device: Device, - eos_token_id1: u32, - eos_token_id2: u32, generation_config: HunyuanOCRGenerationConfig, model_name: String, } @@ -49,10 +46,12 @@ impl<'a> HunyuanOCRGenerateModel<'a> { let pre_processor = HunyuanVLProcessor::new(path, &device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let hunyuan_vl = HunyuanVLModel::new(vb, cfg.clone())?; let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: HunyuanOCRGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let hunyuan_vl = + HunyuanVLModel::new(vb, cfg.clone(), generation_config.eos_token_id.clone())?; + let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -64,8 +63,6 @@ impl<'a> HunyuanOCRGenerateModel<'a> { pre_processor, hunyuan_vl, device, - eos_token_id1: generation_config.eos_token_id[0] as u32, - eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, model_name, }) @@ -80,51 +77,37 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; let data = self .pre_processor .process_info(&mes, &self.tokenizer, &mes_render)?; - let mut input_ids = data.input_ids; - let mut position_ids = Some(&data.position_ids); - let mut image_mask = Some(&data.image_mask); - let mut pixel_values = data.pixel_values; - let mut image_grid_thw = data.image_grid_thw; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate: Vec = Vec::new(); let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logits = self.hunyuan_vl.forward( - &input_ids, - pixel_values.as_ref(), - image_grid_thw.as_ref(), - image_mask, - position_ids, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - position_ids = None; - image_mask = None; - pixel_values = None; - image_grid_thw = None; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.hunyuan_vl.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let input_ids = data.input_ids; + let mut ctx = GenerationContext::new( + temperature.into(), + top_p.into(), + top_k.into(), + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + let data_vec = vec![ + data.pixel_values, + data.image_grid_thw, + data.image_mask.into(), + data.position_ids.into(), + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.hunyuan_vl, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -144,70 +127,34 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> { let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; let data = self .pre_processor .process_info(&mes, &self.tokenizer, &mes_render)?; - let mut seqlen_offset = 0; let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut error_tokens = Vec::new(); - let mut input_ids = data.input_ids; - let mut position_ids = Some(&data.position_ids); - let mut image_mask = Some(&data.image_mask); - let mut pixel_values = data.pixel_values; - let mut image_grid_thw = data.image_grid_thw; - let mut seq_len = input_ids.dim(1)?; - for _ in 0..sample_len { - let logits = self.hunyuan_vl.forward( - &input_ids, - pixel_values.as_ref(), - image_grid_thw.as_ref(), - image_mask, - position_ids, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - position_ids = None; - image_mask = None; - pixel_values = None; - image_grid_thw = None; - continue; - } - error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - position_ids = None; - image_mask = None; - pixel_values = None; - image_grid_thw = None; - } - self.hunyuan_vl.clear_kv_cache(); - }; + let input_ids = data.input_ids; + let data_vec = vec![ + data.pixel_values, + data.image_grid_thw, + data.image_mask.into(), + data.position_ids.into(), + ]; + let data = MultiModalData::new(data_vec); + let stream = generate_stream_generic( + &mut self.hunyuan_vl, + &self.tokenizer, + input_ids, + data, + temperature.into(), + top_p.into(), + top_k.into(), + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/hunyuan_ocr/model.rs b/src/models/hunyuan_ocr/model.rs index e98836a..2585618 100644 --- a/src/models/hunyuan_ocr/model.rs +++ b/src/models/hunyuan_ocr/model.rs @@ -7,14 +7,19 @@ use candle_nn::{ use crate::{ models::{ - common::modules::{ - GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d, + common::{ + InferenceModel, + modules::{ + GateUpDownMLP, NaiveAttnTwoLinearMLPBlock, eager_attention_forward, get_conv2d, + }, }, hunyuan_ocr::config::{HunYuanVLConfig, HunYuanVLVisionConfig}, }, position_embed::rope::{RoPE, apply_rotary_pos_emb, get_xd_cos_sin}, - utils::interpolate::interpolate_bilinear, - utils::tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor}, + utils::{ + interpolate::interpolate_bilinear, + tensor_utils::{masked_scatter_dim0, prepare_causal_attention_mask, split_tensor}, + }, }; pub struct HunYuanVisionPatchEmbed { @@ -538,10 +543,11 @@ pub struct HunyuanVLModel { vit: HunYuanVisionTransformer, model: HunYuanVLTextModel, lm_head: Linear, + stop_token_ids: Vec, } impl HunyuanVLModel { - pub fn new(vb: VarBuilder, config: HunYuanVLConfig) -> Result { + pub fn new(vb: VarBuilder, config: HunYuanVLConfig, eos_ids: Vec) -> Result { let vit = HunYuanVisionTransformer::new(vb.pp("vit"), &config.vision_config)?; let model = HunYuanVLTextModel::new(vb.pp("model"), &config)?; let lm_head = Linear::new(model.embed_tokens.embeddings().clone(), None); @@ -550,6 +556,7 @@ impl HunyuanVLModel { vit, model, lm_head, + stop_token_ids: eos_ids, }) } pub fn forward( @@ -582,3 +589,42 @@ impl HunyuanVLModel { self.model.clear_kv_cache(); } } + +impl InferenceModel for HunyuanVLModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 4 { + return Err(anyhow::anyhow!( + "HunyuanVL process data error, must have pixel_values, image_grid_thw, image_mask, position_ids" + )); + } + let pixel_values = &data.data_vec[0]; + let image_grid_thw = &data.data_vec[1]; + let image_mask = &data.data_vec[2]; + let position_ids = &data.data_vec[3]; + self.forward( + input_ids, + pixel_values.as_ref(), + image_grid_thw.as_ref(), + image_mask.as_ref(), + position_ids.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/lfm2/generate.rs b/src/models/lfm2/generate.rs index 202fd7b..d81d794 100644 --- a/src/models/lfm2/generate.rs +++ b/src/models/lfm2/generate.rs @@ -1,6 +1,8 @@ -use crate::models::common::generate::get_logit_processor; +use crate::models::common::MultiModalData; +use crate::models::common::generate::{ + GenerationContext, generate_generic, generate_stream_generic, +}; use crate::params::chat::{ChatCompletionParameters, ChatCompletionResponse}; -use crate::utils::build_completion_chunk_response; use crate::{ chat_template::ChatTemplate, models::{ @@ -11,19 +13,17 @@ use crate::{ }, }, tokenizer::TokenizerModel, - utils::{build_completion_response, find_type_files, get_device, get_dtype}, + utils::{find_type_files, get_device, get_dtype}, }; use anyhow::Result; -use candle_core::{DType, Device, Tensor}; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; pub struct Lfm2GenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, device: Device, model: Lfm2Model, - eos_token_id: u32, model_name: String, } impl<'a> Lfm2GenerateModel<'a> { @@ -45,8 +45,8 @@ impl<'a> Lfm2GenerateModel<'a> { }; let dtype = get_dtype(dtype, &cfg_dtype); let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? }; - let model = Lfm2Model::new(vb, &cfg)?; - let eos_token_id = gen_cfg.eos_token_id; + let eos_ids = vec![gen_cfg.eos_token_id]; + let model = Lfm2Model::new(vb, &cfg, eos_ids)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -57,7 +57,6 @@ impl<'a> Lfm2GenerateModel<'a> { tokenizer, device, model, - eos_token_id, model_name, }) } @@ -66,40 +65,28 @@ impl<'a> Lfm2GenerateModel<'a> { impl<'a> GenerateModel for Lfm2GenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut logits = get_logit_processor( + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let sample_len = mes.max_tokens.unwrap_or(1024); + let seed = mes.seed.unwrap_or(34562) as u64; + let mut ctx = GenerationContext::new( mes.temperature, mes.top_p, None, - mes.seed.unwrap_or(34562) as u64, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), ); - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate = vec![]; - let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logit = self.model.forward(&input_ids, seqlen_offset)?; - let logit = logit.squeeze(0)?.squeeze(0)?; - let next_token = logits.sample(&logit)?; - generate.push(next_token); - if next_token == self.eos_token_id { - break; - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - } - self.model.clear_cache(); - let completion_tokens = generate.len() as u32; - let decode = self.tokenizer.token_decode(generate)?; - let mes = build_completion_response( - decode, + + let data = MultiModalData::new(vec![]); + generate_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + &mut ctx, &self.model_name, - Some(completion_tokens), - Some(prompt_tokens), - ); - Ok(mes) + ) } fn generate_stream( @@ -115,50 +102,24 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> { >, > { let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut logits = get_logit_processor( + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let sample_len = mes.max_tokens.unwrap_or(1024); + let data = MultiModalData::new(vec![]); + let seed = mes.seed.unwrap_or(34562) as u64; + let stream = generate_stream_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, mes.temperature, mes.top_p, None, - mes.seed.unwrap_or(34562) as u64, - ); - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; - let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut err_tokens = vec![]; - for _ in 0..sample_len { - let logit = self.model.forward(&input_ids, seqlen_offset)?; - let logit = logit.squeeze(0)?.squeeze(0)?; - let next_token = logits.sample(&logit)?; - let mut decode_ids = vec![]; - if !err_tokens.is_empty() { - decode_ids.extend_from_slice(&err_tokens); - } - decode_ids.push(next_token); - let decode = self.tokenizer.token_decode(decode_ids)?; - if decode.contains("�") { - err_tokens.push(next_token); - if err_tokens.len() > 3 { - err_tokens.clear(); - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - continue; - } - err_tokens.clear(); - let chunk = build_completion_chunk_response(decode, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.eos_token_id { - break; - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - } - self.model.clear_cache(); - }; + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/lfm2/model.rs b/src/models/lfm2/model.rs index 348b957..e7ad4d3 100644 --- a/src/models/lfm2/model.rs +++ b/src/models/lfm2/model.rs @@ -1,6 +1,9 @@ use crate::{ models::{ - common::modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d}, + common::{ + InferenceModel, + modules::{GateUpDownMLP, QKNormAttention, conv1d_depthwise, get_conv1d}, + }, lfm2::config::Lfm2Config, }, position_embed::rope::RoPE, @@ -279,10 +282,11 @@ impl Lfm2Decoder { pub struct Lfm2Model { model: Lfm2Decoder, lm_head: Linear, + stop_token_ids: Vec, } impl Lfm2Model { - pub fn new(vb: VarBuilder, config: &Lfm2Config) -> Result { + pub fn new(vb: VarBuilder, config: &Lfm2Config, eos_ids: Vec) -> Result { let model = Lfm2Decoder::new(vb.pp("model"), config)?; let lm_head = if let Some(flag) = config.tie_embedding && flag @@ -300,7 +304,11 @@ impl Lfm2Model { Err(_) => Linear::new(model.embed_tokens.embeddings().clone(), None), } }; - Ok(Self { model, lm_head }) + Ok(Self { + model, + lm_head, + stop_token_ids: eos_ids, + }) } pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { @@ -315,3 +323,17 @@ impl Lfm2Model { self.model.clear_cache(); } } + +impl InferenceModel for Lfm2Model { + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/lfm2vl/generate.rs b/src/models/lfm2vl/generate.rs index db25985..e78c6f3 100644 --- a/src/models/lfm2vl/generate.rs +++ b/src/models/lfm2vl/generate.rs @@ -1,9 +1,12 @@ use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionParameters, ChatCompletionResponse}, }; use anyhow::Result; -use candle_core::{DType, Device, Tensor}; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; use crate::{ @@ -14,12 +17,8 @@ use crate::{ lfm2vl::{config::Lfm2VLConfig, model::Lfm2VLModel, processor::Lfm2VLProcessor}, }, tokenizer::TokenizerModel, - utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, - }, + utils::{find_type_files, get_device, get_dtype}, }; -use rocket::async_stream::stream; pub struct Lfm2VLGenerateModel<'a> { chat_template: ChatTemplate<'a>, @@ -27,7 +26,6 @@ pub struct Lfm2VLGenerateModel<'a> { device: Device, model: Lfm2VLModel, processor: Lfm2VLProcessor, - eos_token_id: u32, model_name: String, } impl<'a> Lfm2VLGenerateModel<'a> { @@ -43,9 +41,9 @@ impl<'a> Lfm2VLGenerateModel<'a> { let model_path = find_type_files(path, "safetensors")?; let dtype = get_dtype(dtype, &cfg.dtype); let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_path, dtype, &device)? }; - let model = Lfm2VLModel::new(vb, &cfg)?; + let eos_ids = vec![gen_cfg.eos_token_id]; + let model = Lfm2VLModel::new(vb, &cfg, eos_ids)?; let processor = Lfm2VLProcessor::new(path, dtype, &device)?; - let eos_token_id = gen_cfg.eos_token_id; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -57,7 +55,6 @@ impl<'a> Lfm2VLGenerateModel<'a> { device, model, processor, - eos_token_id, model_name, }) } @@ -66,54 +63,35 @@ impl<'a> Lfm2VLGenerateModel<'a> { impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut logits = get_logit_processor( + let (pixel_values, pixel_attention_mask, spatial_shapes, text) = + self.processor.process_info(&mes, &mes_render)?; + let input_ids = self.tokenizer.text_encode(text, &self.device)?; + let seed = mes.seed.unwrap_or(34562) as u64; + let sample_len = mes.max_tokens.unwrap_or(1024); + let mut ctx = GenerationContext::new( mes.temperature, mes.top_p, None, - mes.seed.unwrap_or(34562) as u64, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), ); - let (pixel_values, pixel_attention_mask, spatial_shapes, text) = - self.processor.process_info(&mes, &mes_render)?; - let mut input_ids = self.tokenizer.text_encode(text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate = vec![]; - let sample_len = mes.max_tokens.unwrap_or(1024); - let mut pixel_values = Some(pixel_values); - let mut pixel_attention_mask = Some(pixel_attention_mask); - let mut spatial_shapes = Some(spatial_shapes); - for _ in 0..sample_len { - let logit = self.model.forward( - &input_ids, - pixel_values.as_ref(), - pixel_attention_mask.as_ref(), - spatial_shapes.as_ref(), - seqlen_offset, - )?; - let logit = logit.squeeze(0)?.squeeze(0)?; - let next_token = logits.sample(&logit)?; - generate.push(next_token); - if next_token == self.eos_token_id { - break; - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - pixel_values = None; - pixel_attention_mask = None; - spatial_shapes = None; - } - self.model.clear_cache(); - let completion_tokens = generate.len() as u32; - let decode = self.tokenizer.token_decode(generate)?; - let mes = build_completion_response( - decode, + + let data_vec = vec![ + pixel_values.into(), + pixel_attention_mask.into(), + spatial_shapes.into(), + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + &mut ctx, &self.model_name, - Some(completion_tokens), - Some(prompt_tokens), - ); - Ok(mes) + ) } fn generate_stream( @@ -129,67 +107,31 @@ impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> { >, > { let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut logits = get_logit_processor( + let (pixel_values, pixel_attention_mask, spatial_shapes, text) = + self.processor.process_info(&mes, &mes_render)?; + let input_ids = self.tokenizer.text_encode(text, &self.device)?; + let sample_len = mes.max_tokens.unwrap_or(1024); + let data_vec = vec![ + pixel_values.into(), + pixel_attention_mask.into(), + spatial_shapes.into(), + ]; + let data = MultiModalData::new(data_vec); + let seed = mes.seed.unwrap_or(34562) as u64; + let stream = generate_stream_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, mes.temperature, mes.top_p, None, - mes.seed.unwrap_or(34562) as u64, - ); - let (pixel_values, pixel_attention_mask, spatial_shapes, text) = - self.processor.process_info(&mes, &mes_render)?; - let mut input_ids = self.tokenizer.text_encode(text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; - let sample_len = mes.max_tokens.unwrap_or(1024); - let mut pixel_values = Some(pixel_values); - let mut pixel_attention_mask = Some(pixel_attention_mask); - let mut spatial_shapes = Some(spatial_shapes); - let stream = stream! { - let mut err_tokens = vec![]; - for _ in 0..sample_len { - let logit = self.model.forward( - &input_ids, - pixel_values.as_ref(), - pixel_attention_mask.as_ref(), - spatial_shapes.as_ref(), - seqlen_offset, - )?; - let logit = logit.squeeze(0)?.squeeze(0)?; - let next_token = logits.sample(&logit)?; - let mut decode_ids = vec![]; - if !err_tokens.is_empty() { - decode_ids.extend_from_slice(&err_tokens); - } - decode_ids.push(next_token); - let decode = self.tokenizer.token_decode(decode_ids)?; - if decode.contains("�") { - err_tokens.push(next_token); - if err_tokens.len() > 3 { - err_tokens.clear(); - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - pixel_values = None; - pixel_attention_mask = None; - spatial_shapes = None; - continue; - } - err_tokens.clear(); - let chunk = build_completion_chunk_response(decode, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.eos_token_id { - break; - } - input_ids = Tensor::new(vec![next_token], &self.device)?.unsqueeze(0)?; - seqlen_offset += seq_len; - seq_len = 1; - pixel_values = None; - pixel_attention_mask = None; - spatial_shapes = None; - } - self.model.clear_cache(); - }; + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/lfm2vl/model.rs b/src/models/lfm2vl/model.rs index 7fcf394..3365a58 100644 --- a/src/models/lfm2vl/model.rs +++ b/src/models/lfm2vl/model.rs @@ -1,6 +1,9 @@ use crate::{ models::{ - common::modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm}, + common::{ + InferenceModel, + modules::{NaiveAttnTwoLinearMLPBlock, get_layer_norm}, + }, lfm2::model::Lfm2Decoder, lfm2vl::config::{Lfm2VLConfig, Lfm2VLVisionConfig}, }, @@ -15,11 +18,7 @@ use candle_nn::{Activation, LayerNorm, Linear, Module, VarBuilder, embedding, li use num::integer::Roots; pub struct Siglip2VisionEmbeddings { - // embed_dim: usize, - // patch_size: usize, patch_embedding: Linear, - // position_embedding_size: usize, - // position_embedding: Embedding, postitional_embeddings: Tensor, } @@ -44,11 +43,7 @@ impl Siglip2VisionEmbeddings { .permute((2, 0, 1))? .unsqueeze(0)?; Ok(Self { - // embed_dim, - // patch_size, patch_embedding, - // position_embedding_size, - // position_embedding, postitional_embeddings, }) } @@ -254,10 +249,11 @@ pub struct Lfm2VLModel { language_model: Lfm2Decoder, lm_head: Linear, img_id: u32, + stop_token_ids: Vec, } impl Lfm2VLModel { - pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig) -> Result { + pub fn new(vb: VarBuilder, cfg: &Lfm2VLConfig, eos_ids: Vec) -> Result { let vb = vb.pp("model"); let vision_tower = Siglip2VisionModel::new(vb.pp("vision_tower"), &cfg.vision_config)?; let multi_modal_projector = @@ -270,6 +266,7 @@ impl Lfm2VLModel { language_model, lm_head, img_id: cfg.image_token_id, + stop_token_ids: eos_ids, }) } @@ -321,3 +318,40 @@ impl Lfm2VLModel { self.language_model.clear_cache(); } } + +impl InferenceModel for Lfm2VLModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 3 { + return Err(anyhow::anyhow!( + "Lfm2VL process data error, must have pixel_values, pixel_attention_mask, spatial_shapes" + )); + } + let pixel_values = &data.data_vec[0]; + let pixel_attention_mask = &data.data_vec[1]; + let spatial_shapes = &data.data_vec[2]; + self.forward( + input_ids, + pixel_values.as_ref(), + pixel_attention_mask.as_ref(), + spatial_shapes.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 58d96a7..d586c34 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -1,20 +1,19 @@ -use crate::models::common::generate::get_logit_processor; +use crate::models::common::MultiModalData; +use crate::models::common::generate::{ + GenerationContext, generate_generic, generate_stream_generic, +}; use crate::params::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; -use anyhow::{Result, anyhow}; -use candle_core::{DType, Device, Tensor}; +use anyhow::Result; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::models::minicpm4::config::MiniCPM4Config; use crate::models::minicpm4::model::MiniCPMModel; // use crate::models::GenerateStream; -use crate::utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, -}; +use crate::utils::{find_type_files, get_device, get_dtype}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; pub struct MiniCPMGenerateModel<'a> { @@ -22,8 +21,6 @@ pub struct MiniCPMGenerateModel<'a> { tokenizer: TokenizerModel, minicpm: MiniCPMModel, device: Device, - endoftext_id: u32, - im_end_id: u32, model_name: String, } @@ -40,7 +37,8 @@ impl<'a> MiniCPMGenerateModel<'a> { let im_end_id = cfg.eos_token_id[1]; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let minicpm = MiniCPMModel::new(vb, cfg)?; + let eos_ids = vec![endoftext_id, im_end_id]; + let minicpm = MiniCPMModel::new(vb, cfg, eos_ids)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -51,8 +49,6 @@ impl<'a> MiniCPMGenerateModel<'a> { tokenizer, minicpm, device: device.clone(), - endoftext_id, - im_end_id, model_name, }) } @@ -60,33 +56,29 @@ impl<'a> MiniCPMGenerateModel<'a> { impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate = Vec::new(); + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let seed = mes.seed.unwrap_or(34562) as u64; let sample_len = mes.max_tokens.unwrap_or(2048); - for _ in 0..sample_len { - let logits = self.minicpm.forward_with_cache(&input_ids, seqlen_offset)?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.endoftext_id || next_token == self.im_end_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.minicpm.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let mut ctx = GenerationContext::new( + mes.temperature, + mes.top_p, + None, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + let data = MultiModalData::new(vec![]); + generate_generic( + &mut self.minicpm, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( &mut self, @@ -100,50 +92,24 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { >, > { let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let data = MultiModalData::new(vec![]); let sample_len = mes.max_tokens.unwrap_or(512); - let stream = stream! { - let mut error_tokens = Vec::new(); - for _ in 0..sample_len { - let logits = self.minicpm.forward_with_cache( - &input_ids, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty(){ - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - continue; - } - error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.endoftext_id || next_token == self.im_end_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - - } - self.minicpm.clear_kv_cache(); - }; + let stream = generate_stream_generic( + &mut self.minicpm, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index 2c43686..a536ada 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -4,7 +4,10 @@ use candle_nn::{Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, rms_n use crate::{ models::{ - common::modules::{GateUpDownMLP, NaiveAttention}, + common::{ + InferenceModel, + modules::{GateUpDownMLP, NaiveAttention}, + }, minicpm4::config::MiniCPM4Config, }, position_embed::rope::compute_default_rope_parameters, @@ -207,10 +210,11 @@ pub struct MiniCPMModel { norm: RmsNorm, rope_emb: MiniCPMLongRoPE, lm_head: Linear, + stop_token_ids: Vec, } impl MiniCPMModel { - pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> Result { + pub fn new(vb: VarBuilder, cfg: MiniCPM4Config, eos_ids: Vec) -> Result { let vb = vb.pp("model"); let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("embed_tokens"))?; let mut layers = Vec::with_capacity(cfg.num_hidden_layers); @@ -229,10 +233,11 @@ impl MiniCPMModel { norm, rope_emb, lm_head, + stop_token_ids: eos_ids, }) } - pub fn forward(&mut self, input_ids: &Tensor, position_id: usize) -> Result { + pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { let (bs, seq_len) = input_ids.dims2()?; let input_embeds = self .embed_tokens @@ -251,7 +256,7 @@ impl MiniCPMModel { } }; - let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; + let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &self.layers { hidden_states = @@ -267,7 +272,11 @@ impl MiniCPMModel { Ok(logits) } - pub fn forward_with_cache(&mut self, input_ids: &Tensor, position_id: usize) -> Result { + pub fn forward_with_cache( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + ) -> Result { let (bs, seq_len) = input_ids.dims2()?; let input_embeds = self .embed_tokens @@ -285,7 +294,7 @@ impl MiniCPMModel { )?) } }; - let (cos, sin) = self.rope_emb.forward(position_id, seq_len)?; + let (cos, sin) = self.rope_emb.forward(seqlen_offset, seq_len)?; let mut hidden_states = input_embeds; for decode_layer in &mut self.layers { hidden_states = decode_layer.forward_with_cache( @@ -311,3 +320,17 @@ impl MiniCPMModel { } } } + +impl InferenceModel for MiniCPMModel { + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward_with_cache(input_ids, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index 604ff6e..5a39a03 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -1,21 +1,20 @@ -use crate::models::common::generate::get_logit_processor; +use crate::models::common::MultiModalData; +use crate::models::common::generate::{ + GenerationContext, generate_generic, generate_stream_generic, +}; use crate::params::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; -use anyhow::{Result, anyhow}; +use anyhow::Result; use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::models::paddleocr_vl::config::{PaddleOCRVLConfig, PaddleOCRVLPreprocessorConfig}; use crate::models::paddleocr_vl::model::PaddleOCRVLModel; use crate::models::paddleocr_vl::processor::PaddleOCRVLProcessor; use crate::utils::tensor_utils::get_equal_mask; -use crate::utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, -}; +use crate::utils::{find_type_files, get_device, get_dtype}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; pub struct PaddleOCRVLGenerateModel<'a> { @@ -25,7 +24,6 @@ pub struct PaddleOCRVLGenerateModel<'a> { paddleocr_vl: PaddleOCRVLModel, cfg: PaddleOCRVLConfig, device: Device, - end_token_id: u32, model_name: String, } @@ -42,10 +40,9 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { let processor_cfg: PaddleOCRVLPreprocessorConfig = serde_json::from_slice(&std::fs::read(processor_cfg_path)?)?; let pre_processor = PaddleOCRVLProcessor::new(processor_cfg, device, dtype)?; - let end_token_id = 2; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb)?; + let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb, vec![2])?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -58,7 +55,6 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { paddleocr_vl, cfg, device: device.clone(), - end_token_id, model_name, }) } @@ -66,53 +62,43 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - let (replace_text, mut pixel_values, mut image_grid_thw) = + let (replace_text, pixel_values, image_grid_thw) = self.pre_processor.process_info(&mes, &mes_render)?; - let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; + let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; - let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)? + let cache_position = Tensor::ones_like(&input_ids.i(0)?)? .to_dtype(candle_core::DType::F64)? .cumsum(D::Minus1)? .to_dtype(candle_core::DType::U32)? .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; - - let mut generate = Vec::new(); + let seed = mes.seed.unwrap_or(34562) as u64; let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logits = self.paddleocr_vl.forward( - &input_ids, - pixel_values.as_ref(), - image_grid_thw.as_ref(), - &image_mask, - Some(&cache_position), - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.end_token_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - image_grid_thw = None; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.paddleocr_vl.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let mut ctx = GenerationContext::new( + mes.temperature, + mes.top_p, + None, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + let data_vec = vec![ + pixel_values, + image_grid_thw, + image_mask.into(), + cache_position.into(), + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.paddleocr_vl, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -126,72 +112,41 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> { + '_, >, > { - let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; let (replace_text, pixel_values, image_grid_thw) = self.pre_processor.process_info(&mes, &mes_render)?; - let mut input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; + let input_ids = self.tokenizer.text_encode(replace_text, &self.device)?; let image_mask = get_equal_mask(&input_ids, self.cfg.image_token_id)?; - let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)? + let cache_position = Tensor::ones_like(&input_ids.i(0)?)? .to_dtype(candle_core::DType::F64)? .cumsum(D::Minus1)? .to_dtype(candle_core::DType::U32)? .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut error_tokens = Vec::new(); - let mut pixel_values = pixel_values.as_ref(); - let mut image_grid_thw = image_grid_thw.as_ref(); - for _ in 0..sample_len { - let logits = self.paddleocr_vl.forward( - &input_ids, - pixel_values, - image_grid_thw, - &image_mask, - Some(&cache_position), - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - image_grid_thw = None; - continue; - } - error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.end_token_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - image_grid_thw = None; - } - self.paddleocr_vl.clear_kv_cache(); - }; + let data_vec = vec![ + pixel_values, + image_grid_thw, + image_mask.into(), + cache_position.into(), + ]; + let data = MultiModalData::new(data_vec); + let seed = mes.seed.unwrap_or(34562) as u64; + let stream = generate_stream_generic( + &mut self.paddleocr_vl, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/paddleocr_vl/model.rs b/src/models/paddleocr_vl/model.rs index 2305b2c..76eefae 100644 --- a/src/models/paddleocr_vl/model.rs +++ b/src/models/paddleocr_vl/model.rs @@ -8,18 +8,23 @@ use num::integer::Roots; use crate::{ models::{ - common::modules::{ - NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm, + common::{ + InferenceModel, + modules::{ + NaiveAttnGateUpDownMLPBlock, NaiveAttnTwoLinearMLPBlock, get_conv2d, get_layer_norm, + }, }, paddleocr_vl::config::{ PaddleOCRVLConfig, PaddleOCRVLRopeScalingConfig, PaddleOCRVLVisionConfig, }, }, position_embed::rope::{Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding}, - utils::interpolate::interpolate_bilinear, - utils::tensor_utils::{ - get_vision_next_indices, masked_scatter_dim0, nonzero_index, prepare_causal_attention_mask, - zero_index, + utils::{ + interpolate::interpolate_bilinear, + tensor_utils::{ + get_vision_next_indices, masked_scatter_dim0, nonzero_index, + prepare_causal_attention_mask, zero_index, + }, }, }; @@ -412,10 +417,11 @@ pub struct PaddleOCRVLModel { pub cfg: PaddleOCRVLConfig, lm_head: Linear, rope_deltas: Option, + stop_token_ids: Vec, } impl PaddleOCRVLModel { - pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder) -> Result { + pub fn new(cfg: PaddleOCRVLConfig, vb: VarBuilder, eos_ids: Vec) -> Result { let mlp_ar = Projector::new(vb.pp("mlp_AR"), &cfg)?; let visual = SiglipVisionModel::new(vb.pp("visual"), &cfg.vision_config)?; let model = Ernie4_5Model::new(vb.pp("model"), &cfg)?; @@ -433,6 +439,7 @@ impl PaddleOCRVLModel { cfg, lm_head, rope_deltas: None, + stop_token_ids: eos_ids, }) } @@ -662,13 +669,14 @@ impl PaddleOCRVLModel { input_ids: &Tensor, pixel_values: Option<&Tensor>, image_grid_thw: Option<&Tensor>, - image_mask: &Tensor, + image_mask: Option<&Tensor>, cache_position: Option<&Tensor>, seqlen_offset: usize, ) -> Result { let mut inputs_embeds = self.model.embed_tokens.forward(input_ids)?; if let Some(pixel_values) = pixel_values && let Some(image_grid_thw) = image_grid_thw + && let Some(image_mask) = image_mask { let pixel_values = pixel_values.unsqueeze(0)?; let mut siglip_position_ids = vec![]; @@ -716,6 +724,15 @@ impl PaddleOCRVLModel { .broadcast_add(rope_deltas)? .contiguous()? .to_dtype(candle_core::DType::U32)? + } else if let Some(rope_deltas) = &self.rope_deltas { + let cache_position = + Tensor::from_vec(vec![seqlen_offset as u32], 1, inputs_embeds.device())?; + cache_position + .i(0)? + .to_dtype(rope_deltas.dtype())? + .broadcast_add(rope_deltas)? + .contiguous()? + .to_dtype(candle_core::DType::U32)? } else { Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())? }; @@ -740,3 +757,42 @@ impl PaddleOCRVLModel { self.model.clear_kv_cache(); } } + +impl InferenceModel for PaddleOCRVLModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 4 { + return Err(anyhow::anyhow!( + "Lfm2VL process data error, must have pixel_values, image_grid_thw, image_mask, cache_position" + )); + } + let pixel_values = &data.data_vec[0]; + let image_grid_thw = &data.data_vec[1]; + let image_mask = &data.data_vec[2]; + let cache_position = &data.data_vec[3]; + self.forward( + input_ids, + pixel_values.as_ref(), + image_grid_thw.as_ref(), + image_mask.as_ref(), + cache_position.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 94531dd..3203c09 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -1,7 +1,12 @@ +use std::time::Instant; + use crate::models::common::generate::get_logit_processor; use crate::params::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; +use crate::utils::response_utils::{ + build_chunk_response_with_usage, build_completion_response_with_time, +}; use anyhow::{Result, anyhow}; use candle_core::{D, DType, Device, IndexOp, Tensor}; use candle_nn::VarBuilder; @@ -10,8 +15,7 @@ use rocket::futures::Stream; use crate::models::qwen2_5vl::config::Qwen2_5VLConfig; use crate::utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, + find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response, }; use crate::{ chat_template::ChatTemplate, @@ -94,7 +98,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let mut generate = Vec::new(); let sample_len = mes.max_tokens.unwrap_or(1024); + let mut prompt_secs = 0.0f64; + let mut completion_secs = 0.0f64; for _ in 0..sample_len { + let i_start = Instant::now(); let logits = self.qwen2_5_vl.forward( &input_ids, pixel_values, @@ -108,6 +115,12 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { )?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; + let i_duration = i_start.elapsed(); + if seqlen_offset == 0 { + prompt_secs += i_duration.as_secs_f64(); + } else { + completion_secs += i_duration.as_secs_f64(); + }; generate.push(next_token); if next_token == self.endoftext_id || next_token == self.im_end_id { break; @@ -124,8 +137,14 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let num_token = generate.len() as u32; let res = self.tokenizer.token_decode(generate)?; self.qwen2_5_vl.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); + let response = build_completion_response_with_time( + res, + &self.model_name, + num_token.into(), + completion_secs.into(), + prompt_tokens.into(), + prompt_secs.into(), + ); Ok(response) } @@ -148,6 +167,10 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { .tokenizer .text_encode(input.replace_text.clone(), &self.device)?; let mut seq_len = input_ids.dim(1)?; + let prompt_tokens = seq_len as u32; + let mut prompt_secs = 0.0f64; + let mut completion_tokens = 0u32; + let mut completion_secs = 0.0f64; let mut seqlen_offset = 0; let mut mask = Tensor::ones_like(&input_ids)?; let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)? @@ -166,6 +189,7 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let mut tool_call_id = None; let mut tool_call_content = String::new(); for _ in 0..sample_len { + let i_start = Instant::now(); let logits = self.qwen2_5_vl.forward( &input_ids, pixel_values, @@ -179,6 +203,13 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { )?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; + completion_tokens += 1; + let i_duration = i_start.elapsed(); + if seqlen_offset == 0 { + prompt_secs += i_duration.as_secs_f64(); + } else { + completion_secs += i_duration.as_secs_f64(); + }; let mut decode_ids = Vec::new(); if !error_tokens.is_empty() { decode_ids.extend_from_slice(&error_tokens); @@ -249,9 +280,8 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { } } } - // let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - // yield Ok(chunk); if next_token == self.endoftext_id || next_token == self.im_end_id { + yield Ok(build_chunk_response_with_usage(&self.model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into())); break; } seqlen_offset += seq_len; diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs index ffe7b46..5c65191 100644 --- a/src/models/qwen3/generate.rs +++ b/src/models/qwen3/generate.rs @@ -1,20 +1,18 @@ -use crate::models::common::generate::get_logit_processor; +use crate::models::common::MultiModalData; +use crate::models::common::generate::{ + GenerationContext, generate_generic, generate_stream_generic, +}; use crate::params::chat::{ ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, }; -use anyhow::{Result, anyhow}; -use candle_core::{DType, Device, Tensor}; +use anyhow::Result; +use candle_core::{DType, Device}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig}; use crate::models::qwen3::model::Qwen3Model; -// use crate::models::GenerateStream; -use crate::utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, -}; +use crate::utils::{find_type_files, get_device, get_dtype}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; pub struct Qwen3GenerateModel<'a> { @@ -22,8 +20,6 @@ pub struct Qwen3GenerateModel<'a> { tokenizer: TokenizerModel, qwen3: Qwen3Model, device: Device, - eos_token_id1: u32, - eos_token_id2: u32, generation_config: Qwen3GenerationConfig, model_name: String, } @@ -39,10 +35,11 @@ impl<'a> Qwen3GenerateModel<'a> { let dtype = get_dtype(dtype, cfg_dtype); let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; - let qwen3 = Qwen3Model::new(&cfg, vb)?; let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let qwen3 = Qwen3Model::new(&cfg, vb, generation_config.eos_token_id.clone())?; + let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -53,8 +50,6 @@ impl<'a> Qwen3GenerateModel<'a> { tokenizer, qwen3, device: device.clone(), - eos_token_id1: generation_config.eos_token_id[0] as u32, - eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, model_name, }) @@ -69,38 +64,28 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> { let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); - let mes_render = self.chat_template.apply_chat_template(&mes)?; - // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking"); - // let mes_render = self - // .chat_template - // .apply_chat_temp_think(&mes, enable_thinking)?; - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut generate = Vec::new(); + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(2048); - for _ in 0..sample_len { - let logits = self.qwen3.forward(Some(&input_ids), None, seqlen_offset)?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.qwen3.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let mut ctx = GenerationContext::new( + temperature.into(), + top_p.into(), + top_k.into(), + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + let data = MultiModalData::new(vec![]); + generate_generic( + &mut self.qwen3, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( &mut self, @@ -119,56 +104,25 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> { let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking"); - // let mes_render = self - // .chat_template - // .apply_chat_temp_think(&mes, enable_thinking)?; - let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; + let in_reasoning = mes_render.ends_with("\n"); + let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; + let data = MultiModalData::new(vec![]); let sample_len = mes.max_tokens.unwrap_or(512); - let stream = stream! { - let mut error_tokens = Vec::new(); - for _ in 0..sample_len { - let logits = self.qwen3.forward( - Some(&input_ids), - None, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty(){ - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - continue; - } - error_tokens.clear(); - let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); - yield Ok(chunk); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - - } - self.qwen3.clear_kv_cache(); - }; + let stream = generate_stream_generic( + &mut self.qwen3, + &self.tokenizer, + input_ids, + data, + temperature.into(), + top_p.into(), + top_k.into(), + seed, + sample_len, + in_reasoning, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs index cb6ee0d..ef82def 100644 --- a/src/models/qwen3/model.rs +++ b/src/models/qwen3/model.rs @@ -6,7 +6,10 @@ use candle_nn::{ use crate::{ models::{ - common::modules::{GateUpDownMLP, QKNormAttention}, + common::{ + InferenceModel, + modules::{GateUpDownMLP, QKNormAttention}, + }, qwen3::config::Qwen3Config, }, position_embed::rope::RoPE, @@ -94,10 +97,11 @@ pub struct Qwen3Model { norm: RmsNorm, rotary_emb: RoPE, lm_head: Linear, + stop_token_ids: Vec, } impl Qwen3Model { - pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { + pub fn new(config: &Qwen3Config, vb: VarBuilder, eos_ids: Vec) -> Result { let vb = vb.pp("model"); let vocab_size = config.vocab_size; let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?; @@ -121,6 +125,7 @@ impl Qwen3Model { norm, rotary_emb, lm_head, + stop_token_ids: eos_ids, }) } pub fn forward( @@ -178,3 +183,17 @@ impl Qwen3Model { } } } + +impl InferenceModel for Qwen3Model { + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids.into(), None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs index 58da587..25dc8d2 100644 --- a/src/models/qwen3_5/generate.rs +++ b/src/models/qwen3_5/generate.rs @@ -1,11 +1,13 @@ use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; use anyhow::{Result, anyhow}; -use candle_core::{DType, Device, Tensor, quantized::gguf_file}; +use candle_core::{DType, Device, quantized::gguf_file}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::{ @@ -17,10 +19,7 @@ use crate::{ qwen3vl::processor::Qwen3VLProcessor, }, tokenizer::TokenizerModel, - utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, - }, + utils::{find_type_files, get_device, get_dtype}, }; pub struct Qwen3_5GenerateModel<'a> { @@ -29,10 +28,9 @@ pub struct Qwen3_5GenerateModel<'a> { pre_processor: Option, qwen3_5: Qwen3_5Model, device: Device, - eos_token_id: u32, model_name: String, - repeat_penalty: f32, - repeat_last_n: usize, + // repeat_penalty: f32, // TODO + // repeat_last_n: usize, } impl<'a> Qwen3_5GenerateModel<'a> { @@ -51,8 +49,8 @@ impl<'a> Qwen3_5GenerateModel<'a> { let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let eos_token_id = cfg.text_config.eos_token_id; - let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg)?; + let eos_ids = vec![cfg.text_config.eos_token_id]; + let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg, eos_ids)?; Ok(Self { chat_template, @@ -60,10 +58,9 @@ impl<'a> Qwen3_5GenerateModel<'a> { pre_processor: Some(pre_processor), qwen3_5, device, - eos_token_id, model_name: model_name.to_string(), - repeat_penalty: 1.01, - repeat_last_n: 64, + // repeat_penalty: 1.01, + // repeat_last_n: 64, }) } @@ -105,7 +102,9 @@ impl<'a> Qwen3_5GenerateModel<'a> { let eos_token_id = model_gguf .get_matedata("tokenizer.ggml.eos_token_id")? .to_u32()?; - let qwen3_5 = Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device)?; + let eos_ids = vec![eos_token_id]; + let qwen3_5 = + Qwen3_5Model::new_from_gguf(&mut model_gguf, mmproj_gguf.as_mut(), &device, eos_ids)?; let stem = std::path::Path::new(model_file) .file_stem() // 获取文件名主干(不含扩展名) .and_then(|s| s.to_str()) @@ -116,11 +115,9 @@ impl<'a> Qwen3_5GenerateModel<'a> { pre_processor, qwen3_5, device, - // eos_token_id: 248044, - eos_token_id, model_name: stem.to_string(), - repeat_penalty: 1.1, - repeat_last_n: 64, + // repeat_penalty: 1.1, + // repeat_last_n: 64, }) } } @@ -130,9 +127,6 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { let seed = mes.seed.unwrap_or(32768) as u64; let temperature = mes.temperature.unwrap_or(0.4); let top_p = mes.top_p.unwrap_or(0.95); - let mut logit_processor = - get_logit_processor(temperature.into(), top_p.into(), Some(20), seed); - // let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); 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 { @@ -147,57 +141,32 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { } else { (mes_render, None, None, None, None) }; - let mut input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut pixel_values = pixel_values.as_ref(); - let image_grid_thw = image_grid_thw.as_ref(); - let mut pixel_values_video = pixel_values_video.as_ref(); - let video_grid_thw = video_grid_thw.as_ref(); - let mut generate = Vec::new(); + let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logits = self.qwen3_5.forward( - &input_ids, - pixel_values, - image_grid_thw, - pixel_values_video, - video_grid_thw, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let logits = if self.repeat_penalty == 1. { - logits - } else { - let start_at = generate.len().saturating_sub(self.repeat_last_n); - candle_transformers::utils::apply_repeat_penalty( - &logits, - self.repeat_penalty, - &generate[start_at..], - )? - }; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.eos_token_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - pixel_values = None; - pixel_values_video = None; - } - let completion_tokens = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.qwen3_5.clear_cache(); - let response = build_completion_response( - res, - &self.model_name, - Some(completion_tokens), - Some(prompt_tokens), + let mut ctx = GenerationContext::new( + temperature.into(), + top_p.into(), + Some(20), + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), ); - Ok(response) + let data_vec = vec![ + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.qwen3_5, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -211,10 +180,8 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { + '_, >, > { - let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - // let input = self.pre_processor.process_info(&mes, &mes_render)?; + let in_reasoning = mes_render.ends_with("\n"); 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)?; @@ -228,117 +195,30 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { } else { (mes_render, None, None, None, None) }; - let mut input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; + let input_ids = self.tokenizer.text_encode(mes_text, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut error_tokens = Vec::new(); - let mut pixel_values = pixel_values.as_ref(); - let image_grid_thw = image_grid_thw.as_ref(); - let mut pixel_values_video = pixel_values_video.as_ref(); - let video_grid_thw = video_grid_thw.as_ref(); - let mut tool_call_id = None; - let mut tool_call_content = String::new(); - let mut generate = Vec::new(); - for _ in 0..sample_len { - let logits = self.qwen3_5.forward( - &input_ids, - pixel_values, - image_grid_thw, - pixel_values_video, - video_grid_thw, - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let logits = if self.repeat_penalty == 1. { - logits - } else { - let start_at = generate.len().saturating_sub(self.repeat_last_n); - candle_transformers::utils::apply_repeat_penalty( - &logits, - self.repeat_penalty, - &generate[start_at..], - )? - }; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } - error_tokens.clear(); - // 处理特殊标记和工具调用 - match decoded_token.as_str() { - "" => { - // 开始工具调用 - tool_call_id = Some(uuid::Uuid::new_v4().to_string()); - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } - "" => { - // 结束工具调用 - let chunk = build_completion_chunk_response( - decoded_token, - &self.model_name, - tool_call_id.clone(), - Some(tool_call_content.clone()) - ); - tool_call_id = None; - tool_call_content = String::new(); - yield Ok(chunk); - } - _ => { - if tool_call_id.is_some() { - // 在工具调用过程中,收集工具调用内容 - tool_call_content.push_str(&decoded_token); - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } else { - // 正常文本输出 - let chunk = build_completion_chunk_response( - decoded_token, - &self.model_name, - None, - None - ); - yield Ok(chunk); - } - } - } - if next_token == self.eos_token_id { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - pixel_values = None; - pixel_values_video = None; - } - self.qwen3_5.clear_cache(); - }; + 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; + let stream = generate_stream_generic( + &mut self.qwen3_5, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + seed, + sample_len, + in_reasoning, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/qwen3_5/model.rs b/src/models/qwen3_5/model.rs index 4d08a1a..7b0ffa9 100644 --- a/src/models/qwen3_5/model.rs +++ b/src/models/qwen3_5/model.rs @@ -10,6 +10,7 @@ use candle_nn::{ use crate::{ models::{ common::{ + InferenceModel, gguf::{GateUpDownMLPGguf, Gguf, ProjKind, QuantizedLinear}, modules::{conv1d_depthwise, eager_attention_forward, get_conv1d, softplus}, }, @@ -1043,10 +1044,11 @@ pub struct Qwen3_5Model { language_model: Qwen3_5TextModel, lm_head: ProjKind, rope_deltas: Option, + stop_token_ids: Vec, } impl Qwen3_5Model { - pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config) -> Result { + pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config, eos_ids: Vec) -> Result { let vb_m = vb.pp("model"); let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?; let language_model = @@ -1069,6 +1071,7 @@ impl Qwen3_5Model { language_model, lm_head: ProjKind::LinearProj(lm_head), rope_deltas: None, + stop_token_ids: eos_ids, }) } @@ -1076,6 +1079,7 @@ impl Qwen3_5Model { gguf: &mut Gguf, mmproj_gguf: Option<&mut Gguf>, device: &Device, + eos_ids: Vec, ) -> Result { let spatial_merge_size = 2usize; let image_token_id = 248056u32; @@ -1102,6 +1106,7 @@ impl Qwen3_5Model { language_model, lm_head: ProjKind::QuantizedProj(QuantizedLinear::new(lm_head, None)), rope_deltas: None, + stop_token_ids: eos_ids, }) } @@ -1434,6 +1439,46 @@ impl Qwen3_5Model { } pub fn clear_cache(&mut self) { + self.rope_deltas = None; self.language_model.clear_cache(); } } + +impl InferenceModel for Qwen3_5Model { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 4 { + return Err(anyhow::anyhow!( + "Lfm2VL process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw" + )); + } + let pixel_values = &data.data_vec[0]; + let image_grid_thw = &data.data_vec[1]; + let pixel_values_video = &data.data_vec[2]; + let video_grid_thw = &data.data_vec[3]; + self.forward( + input_ids, + pixel_values.as_ref(), + image_grid_thw.as_ref(), + pixel_values_video.as_ref(), + video_grid_thw.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/qwen3_asr/config.rs b/src/models/qwen3_asr/config.rs index 853f186..579124a 100644 --- a/src/models/qwen3_asr/config.rs +++ b/src/models/qwen3_asr/config.rs @@ -202,7 +202,7 @@ pub struct Qwen3ASRRopeScaling { #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct Qwen3ASRGenerationConfig { pub do_sample: bool, - pub eos_token_id: Vec, + pub eos_token_id: Vec, pub pad_token_id: usize, pub temperature: f32, } diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index d7a881e..019ec0a 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -1,6 +1,9 @@ +use std::time::Instant; + use crate::{ models::common::generate::get_logit_processor, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, + utils::response_utils::{build_chunk_response_with_usage, build_completion_response_with_time}, }; use anyhow::{Result, anyhow}; use candle_core::{DType, Device, Tensor}; @@ -21,8 +24,7 @@ use crate::{ }, tokenizer::TokenizerModel, utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, + find_type_files, get_device, get_dtype, response_utils::build_completion_chunk_response, }, }; @@ -92,6 +94,8 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { let sample_len = mes.max_tokens.unwrap_or(1024); let mut generate = Vec::new(); let mut prompt_tokens = 0u32; + let mut prompt_secs = 0.0f64; + let mut completion_secs = 0.0f64; for data in audio_datas.iter() { let mut input_ids = data.input_ids.clone(); let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?); @@ -99,11 +103,18 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { prompt_tokens += seq_len as u32; let mut seqlen_offset = 0; for _ in 0..sample_len { + let i_start = Instant::now(); let logits = self.qwen3_asr .forward(&input_ids, seqlen_offset, input_features.as_ref())?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; + let i_duration = i_start.elapsed(); + if seqlen_offset == 0 { + prompt_secs += i_duration.as_secs_f64(); + } else { + completion_secs += i_duration.as_secs_f64(); + }; generate.push(next_token); if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { break; @@ -117,8 +128,14 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { } let num_token = generate.len() as u32; let res = self.tokenizer.token_decode(generate)?; - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); + let response = build_completion_response_with_time( + res, + &self.model_name, + num_token.into(), + completion_secs.into(), + prompt_tokens.into(), + prompt_secs.into(), + ); Ok(response) } @@ -145,17 +162,30 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { let sample_len = mes.max_tokens.unwrap_or(1024); let stream = stream! { let mut error_tokens = Vec::new(); + let mut prompt_tokens = 0u32; + let mut completion_tokens = 0u32; + let mut prompt_secs = 0.0f64; + let mut completion_secs = 0.0f64; for data in audio_datas.iter() { let mut input_ids = data.input_ids.clone(); let mut input_features = Some(data.input_features.clone().to_dtype(self.dtype)?); let mut seq_len = input_ids.dim(1)?; + prompt_tokens += seq_len as u32; let mut seqlen_offset = 0; for _ in 0..sample_len { + let i_start = Instant::now(); let logits = self.qwen3_asr .forward(&input_ids, seqlen_offset, input_features.as_ref())?; let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; let next_token = logit_processor.sample(&logits)?; + completion_tokens += 1; + let i_duration = i_start.elapsed(); + if seqlen_offset == 0 { + prompt_secs += i_duration.as_secs_f64(); + } else { + completion_secs += i_duration.as_secs_f64(); + }; let mut decode_ids = Vec::new(); if !error_tokens.is_empty() { decode_ids.extend_from_slice(&error_tokens); @@ -177,6 +207,7 @@ impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> { let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None); yield Ok(chunk); if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + yield Ok(build_chunk_response_with_usage(&self.model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into())); break; } seqlen_offset += seq_len; diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 16760cd..f4b91ca 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -1,11 +1,13 @@ use crate::{ - models::common::generate::get_logit_processor, + models::common::{ + MultiModalData, + generate::{GenerationContext, generate_generic, generate_stream_generic}, + }, params::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}, }; -use anyhow::{Result, anyhow}; +use anyhow::Result; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; -use rocket::async_stream::stream; use rocket::futures::Stream; use crate::{ @@ -16,10 +18,7 @@ use crate::{ qwen3vl::{config::Qwen3VLConfig, model::Qwen3VLModel, processor::Qwen3VLProcessor}, }, tokenizer::TokenizerModel, - utils::{ - build_completion_chunk_response, build_completion_response, find_type_files, get_device, - get_dtype, - }, + utils::{find_type_files, get_device, get_dtype}, }; pub struct Qwen3VLGenerateModel<'a> { @@ -28,8 +27,6 @@ pub struct Qwen3VLGenerateModel<'a> { pre_processor: Qwen3VLProcessor, qwen3_vl: Qwen3VLModel, device: Device, - eos_token_id1: u32, - eos_token_id2: u32, generation_config: Qwen3GenerationConfig, model_name: String, } @@ -46,10 +43,11 @@ impl<'a> Qwen3VLGenerateModel<'a> { let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?; let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; - let qwen3_vl = Qwen3VLModel::new(cfg, vb)?; let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let qwen3_vl = Qwen3VLModel::new(cfg, vb, generation_config.eos_token_id.clone())?; + let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) @@ -61,8 +59,6 @@ impl<'a> Qwen3VLGenerateModel<'a> { pre_processor, qwen3_vl, device, - eos_token_id1: generation_config.eos_token_id[0] as u32, - eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, model_name, }) @@ -77,56 +73,39 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; - // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking"); - // let mes_render = self - // .chat_template - // .apply_chat_temp_think(&mes, enable_thinking)?; let input = self.pre_processor.process_info(&mes, &mes_render)?; - let mut input_ids = self + let input_ids = self .tokenizer .text_encode(input.replace_text.clone(), &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let prompt_tokens = seq_len as u32; - let mut seqlen_offset = 0; - let mut pixel_values = input.pixel_values.as_ref(); - let image_grid_thw = input.image_grid_thw.as_ref(); - let mut pixel_values_video = input.pixel_values_video.as_ref(); - let video_grid_thw = input.video_grid_thw.as_ref(); - let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; - let mut generate = Vec::new(); + let seq_len = input_ids.dim(1)?; let sample_len = mes.max_tokens.unwrap_or(1024); - for _ in 0..sample_len { - let logits = self.qwen3_vl.forward( - &input_ids, - pixel_values, - image_grid_thw, - pixel_values_video, - video_grid_thw, - Some(&cache_position), - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - generate.push(next_token); - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - pixel_values_video = None; - } - let num_token = generate.len() as u32; - let res = self.tokenizer.token_decode(generate)?; - self.qwen3_vl.clear_kv_cache(); - let response = - build_completion_response(res, &self.model_name, Some(num_token), Some(prompt_tokens)); - Ok(response) + let mut ctx = GenerationContext::new( + temperature.into(), + top_p.into(), + top_k.into(), + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; + let data_vec = vec![ + input.pixel_values, + input.image_grid_thw, + input.pixel_values_video, + input.video_grid_thw, + cache_position.into(), + ]; + let data = MultiModalData::new(data_vec); + generate_generic( + &mut self.qwen3_vl, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) } fn generate_stream( @@ -145,119 +124,141 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { .unwrap_or(self.generation_config.temperature); let top_p = mes.top_p.unwrap_or(self.generation_config.top_p); let top_k = self.generation_config.top_k; - let seed = mes.seed.unwrap_or(34562) as u64; - let mut logit_processor = - get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); let mes_render = self.chat_template.apply_chat_template(&mes)?; + let in_reasoning = mes_render.ends_with("\n"); let input = self.pre_processor.process_info(&mes, &mes_render)?; - let mut input_ids = self + let input_ids = self .tokenizer .text_encode(input.replace_text.clone(), &self.device)?; - let mut seq_len = input_ids.dim(1)?; - let mut seqlen_offset = 0; - let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; + let seq_len = input_ids.dim(1)?; + let cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; let sample_len = mes.max_tokens.unwrap_or(1024); - let stream = stream! { - let mut error_tokens = Vec::new(); - let mut pixel_values = input.pixel_values.as_ref(); - let image_grid_thw = input.image_grid_thw.as_ref(); - let mut pixel_values_video = input.pixel_values_video.as_ref(); - let video_grid_thw = input.video_grid_thw.as_ref(); - let mut tool_call_id = None; - let mut tool_call_content = String::new(); - for _ in 0..sample_len { - let logits = self.qwen3_vl.forward( - &input_ids, - pixel_values, - image_grid_thw, - pixel_values_video, - video_grid_thw, - Some(&cache_position), - seqlen_offset, - )?; - let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; - let next_token = logit_processor.sample(&logits)?; - let mut decode_ids = Vec::new(); - if !error_tokens.is_empty() { - decode_ids.extend_from_slice(&error_tokens); - } - decode_ids.push(next_token); - let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; - if decoded_token.contains("�") { - error_tokens.push(next_token); - if error_tokens.len() > 3 { - error_tokens.clear(); - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } - error_tokens.clear(); - - // 处理特殊标记和工具调用 - match decoded_token.as_str() { - "" => { - // 开始工具调用 - tool_call_id = Some(uuid::Uuid::new_v4().to_string()); - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } - "" => { - // 结束工具调用 - let chunk = build_completion_chunk_response( - decoded_token, - &self.model_name, - tool_call_id.clone(), - Some(tool_call_content.clone()) - ); - tool_call_id = None; - tool_call_content = String::new(); - yield Ok(chunk); - } - _ => { - if tool_call_id.is_some() { - // 在工具调用过程中,收集工具调用内容 - tool_call_content.push_str(&decoded_token); - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - pixel_values_video = None; - continue; - } else { - // 正常文本输出 - let chunk = build_completion_chunk_response( - decoded_token, - &self.model_name, - None, - None - ); - yield Ok(chunk); - } - } - } - if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { - break; - } - seqlen_offset += seq_len; - seq_len = 1; - input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; - cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; - pixel_values = None; - pixel_values_video = None; - } - self.qwen3_vl.clear_kv_cache(); - }; + let data_vec = vec![ + input.pixel_values, + input.image_grid_thw, + input.pixel_values_video, + input.video_grid_thw, + cache_position.into(), + ]; + let data = MultiModalData::new(data_vec); + let seed = mes.seed.unwrap_or(34562) as u64; + let stream = generate_stream_generic( + &mut self.qwen3_vl, + &self.tokenizer, + input_ids, + data, + temperature.into(), + top_p.into(), + top_k.into(), + seed, + sample_len, + in_reasoning, + &self.device, + &self.model_name, + )?; Ok(Box::new(Box::pin(stream))) + // let sample_len = mes.max_tokens.unwrap_or(1024); + // let stream = stream! { + // let mut error_tokens = Vec::new(); + // let mut pixel_values = input.pixel_values.as_ref(); + // let image_grid_thw = input.image_grid_thw.as_ref(); + // let mut pixel_values_video = input.pixel_values_video.as_ref(); + // let video_grid_thw = input.video_grid_thw.as_ref(); + // let mut tool_call_id = None; + // let mut tool_call_content = String::new(); + // for _ in 0..sample_len { + // let logits = self.qwen3_vl.forward( + // &input_ids, + // pixel_values, + // image_grid_thw, + // pixel_values_video, + // video_grid_thw, + // Some(&cache_position), + // seqlen_offset, + // )?; + // let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + // let next_token = logit_processor.sample(&logits)?; + // let mut decode_ids = Vec::new(); + // if !error_tokens.is_empty() { + // decode_ids.extend_from_slice(&error_tokens); + // } + // decode_ids.push(next_token); + // let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; + // if decoded_token.contains("�") { + // error_tokens.push(next_token); + // if error_tokens.len() > 3 { + // error_tokens.clear(); + // } + // seqlen_offset += seq_len; + // seq_len = 1; + // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + // pixel_values = None; + // pixel_values_video = None; + // continue; + // } + // error_tokens.clear(); + + // // 处理特殊标记和工具调用 + // match decoded_token.as_str() { + // "" => { + // // 开始工具调用 + // tool_call_id = Some(uuid::Uuid::new_v4().to_string()); + // seqlen_offset += seq_len; + // seq_len = 1; + // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + // pixel_values = None; + // pixel_values_video = None; + // continue; + // } + // "" => { + // // 结束工具调用 + // let chunk = build_completion_chunk_response( + // decoded_token, + // &self.model_name, + // tool_call_id.clone(), + // Some(tool_call_content.clone()) + // ); + // tool_call_id = None; + // tool_call_content = String::new(); + // yield Ok(chunk); + // } + // _ => { + // if tool_call_id.is_some() { + // // 在工具调用过程中,收集工具调用内容 + // tool_call_content.push_str(&decoded_token); + // seqlen_offset += seq_len; + // seq_len = 1; + // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + // pixel_values = None; + // pixel_values_video = None; + // continue; + // } else { + // // 正常文本输出 + // let chunk = build_completion_chunk_response( + // decoded_token, + // &self.model_name, + // None, + // None + // ); + // yield Ok(chunk); + // } + // } + // } + // if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + // break; + // } + // seqlen_offset += seq_len; + // seq_len = 1; + // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + // pixel_values = None; + // pixel_values_video = None; + // } + // self.qwen3_vl.clear_kv_cache(); + // }; + // Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs index c166d18..1e21a7f 100644 --- a/src/models/qwen3vl/model.rs +++ b/src/models/qwen3vl/model.rs @@ -10,6 +10,7 @@ use candle_nn::{ use crate::{ models::{ common::{ + InferenceModel, gguf::{Gguf, ProjKind, TwoLinearMLPGguf}, modules::{eager_attention_forward, get_layer_norm}, }, @@ -839,10 +840,11 @@ pub struct Qwen3VLModel { language_model: Qwen3VLTextModel, lm_head: Linear, rope_deltas: Option, + stop_token_ids: Vec, } impl Qwen3VLModel { - pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result { + pub fn new(config: Qwen3VLConfig, vb: VarBuilder, eos_ids: Vec) -> Result { let vb_m = vb.pp("model"); let config = config.clone(); let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?; @@ -863,6 +865,7 @@ impl Qwen3VLModel { language_model, lm_head, rope_deltas: None, + stop_token_ids: eos_ids, }) } @@ -1240,6 +1243,15 @@ impl Qwen3VLModel { .broadcast_add(rope_deltas)? .contiguous()? .to_dtype(candle_core::DType::U32)? + } else if let Some(rope_deltas) = &self.rope_deltas { + let cache_position = + Tensor::from_vec(vec![seqlen_offset as u32], 1, inputs_embeds.device())?; + cache_position + .i(0)? + .to_dtype(rope_deltas.dtype())? + .broadcast_add(rope_deltas)? + .contiguous()? + .to_dtype(candle_core::DType::U32)? } else { Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())? }; @@ -1265,6 +1277,48 @@ impl Qwen3VLModel { } pub fn clear_kv_cache(&mut self) { + self.rope_deltas = None; self.language_model.clear_kv_cache(); } } + +impl InferenceModel for Qwen3VLModel { + fn forward_initial( + &mut self, + input_ids: &Tensor, + seqlen_offset: usize, + data: crate::models::common::MultiModalData, + ) -> Result { + if data.data_vec.len() != 5 { + return Err(anyhow::anyhow!( + "Qwen3VL process data error, must have pixel_values, image_grid_thw, pixel_values_video, video_grid_thw, cache_position" + )); + } + let pixel_values = &data.data_vec[0]; + let image_grid_thw = &data.data_vec[1]; + let pixel_values_video = &data.data_vec[2]; + let video_grid_thw = &data.data_vec[3]; + let cache_position = &data.data_vec[4]; + self.forward( + input_ids, + pixel_values.as_ref(), + image_grid_thw.as_ref(), + pixel_values_video.as_ref(), + video_grid_thw.as_ref(), + cache_position.as_ref(), + seqlen_offset, + ) + } + + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + self.forward(input_ids, None, None, None, None, None, seqlen_offset) + } + + fn clear_cache(&mut self) { + self.clear_kv_cache(); + } + + fn stop_token_ids(&self) -> Vec { + self.stop_token_ids.clone() + } +} diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index 28db495..a8f25e7 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -14,8 +14,9 @@ use rocket::futures::{Stream, stream}; use crate::{ models::{GenerateModel, rmbg2_0::model::BiRefNet}, utils::{ - build_img_completion_response, find_type_files, get_device, get_dtype, + find_type_files, get_device, get_dtype, img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize}, + response_utils::build_img_completion_response, }, }; diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 6dc4301..b12cb08 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -21,8 +21,8 @@ use crate::{ }, utils::{ audio_utils::{extract_audio_url, get_audio_wav_u8}, - build_audio_completion_response, extract_metadata_value, extract_user_text, - find_type_files, get_device, get_dtype, + extract_metadata_value, extract_user_text, find_type_files, get_device, get_dtype, + response_utils::build_audio_completion_response, }, }; diff --git a/src/params/chat.rs b/src/params/chat.rs index 8747bc0..ea314d4 100644 --- a/src/params/chat.rs +++ b/src/params/chat.rs @@ -70,6 +70,8 @@ pub struct ChatCompletionParameters { /// Developer-defined tags and values used for filtering completions in the dashboard. #[serde(skip_serializing_if = "Option::is_none")] pub metadata: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub enable_thinking: Option, /// Number between -2.0 and 2.0. Positive values penalize new tokens based on their existing frequency in the text so far, /// decreasing the model's likelihood to repeat the same line verbatim. #[serde(skip_serializing_if = "Option::is_none")] diff --git a/src/utils/mod.rs b/src/utils/mod.rs index b930531..ea1dc2f 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1,6 +1,7 @@ pub mod audio_utils; pub mod img_utils; pub mod interpolate; +pub mod response_utils; pub mod tensor_utils; pub mod video_utils; @@ -10,14 +11,8 @@ use std::time::{SystemTime, UNIX_EPOCH}; use std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration}; use crate::models::common::model_mapping::WhichModel; -use crate::params::{ - chat::{ - AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse, - ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart, - ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage, - DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall, - }, - shared::{FinishReason, Usage}, +use crate::params::chat::{ + ChatCompletionParameters, ChatMessage, ChatMessageContent, ChatMessageContentPart, }; use anyhow::{Result, anyhow}; use byteorder::{LittleEndian, ReadBytesExt}; @@ -409,305 +404,6 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 { ceil * factor } -pub fn build_img_completion_response( - base64vec: &Vec, - model_name: &str, -) -> ChatCompletionResponse { - let id = uuid::Uuid::new_v4().to_string(); - let mut response = ChatCompletionResponse { - id: Some(id), - choices: vec![], - // created: chrono::Utc::now().timestamp() as u32, - created: timestamp() as u32, - model: model_name.to_string(), - service_tier: None, - system_fingerprint: None, - object: "chat.completion".to_string(), - usage: None, - }; - let mut conten_part_vec = vec![]; - for img_bas64 in base64vec { - let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64; - let part = ChatMessageContentPart::Image(ChatMessageImageContentPart { - r#type: "image".to_string(), - image_url: ImageUrlType { - url: img_base64_prefix, - detail: None, - }, - }); - conten_part_vec.push(part); - } - let choice = ChatCompletionChoice { - index: 0, - message: ChatMessage::Assistant { - content: Some(ChatMessageContent::ContentPart(conten_part_vec)), - reasoning_content: None, - refusal: None, - name: None, - audio: None, - tool_calls: None, - }, - finish_reason: Some(FinishReason::StopSequenceReached), - logprobs: None, - }; - response.choices.push(choice); - response -} - -pub fn build_audio_completion_response( - base64_audio: &String, - model_name: &str, -) -> ChatCompletionResponse { - let id = uuid::Uuid::new_v4().to_string(); - let mut response = ChatCompletionResponse { - id: Some(id), - choices: vec![], - created: timestamp() as u32, - model: model_name.to_string(), - service_tier: None, - system_fingerprint: None, - object: "chat.completion".to_string(), - usage: None, - }; - - let base64_audio = format!("data:audio/wav;base64,{}", base64_audio); - let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart { - r#type: "audio".to_string(), - audio_url: AudioUrlType { - url: base64_audio.to_string(), - }, - })]; - let choice = ChatCompletionChoice { - index: 0, - message: ChatMessage::Assistant { - content: Some(ChatMessageContent::ContentPart(conten_part_vec)), - reasoning_content: None, - refusal: None, - name: None, - audio: None, - tool_calls: None, - }, - finish_reason: Some(FinishReason::StopSequenceReached), - logprobs: None, - }; - response.choices.push(choice); - response -} - -fn build_response(res: String, model_name: &str, usage: Option) -> ChatCompletionResponse { - let id = uuid::Uuid::new_v4().to_string(); - let mut response = ChatCompletionResponse { - id: Some(id), - choices: vec![], - created: timestamp() as u32, - model: model_name.to_string(), - service_tier: None, - system_fingerprint: None, - object: "chat.completion".to_string(), - usage, - }; - let choice = if res.contains("") { - let mes: Vec<&str> = res.split("").collect(); - let content = mes[0].to_string(); - let mut tool_vec = Vec::new(); - for (i, m) in mes.iter().enumerate().skip(1) { - let tool_mes = m.replace("", ""); - let function = match serde_json::from_str::(&tool_mes) { - Ok(json_value) => { - let name = json_value - .get("name") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()) - .unwrap_or_default(); - - let arguments = json_value - .get("arguments") - .map(|v| v.to_string()) - .unwrap_or_default(); - - Function { name, arguments } - } - Err(_) => Function { - name: "".to_string(), - arguments: "".to_string(), - }, - }; - let tool_call = ToolCall { - id: (i - 1).to_string(), - r#type: "function".to_string(), - function, - }; - tool_vec.push(tool_call); - } - ChatCompletionChoice { - index: 0, - message: ChatMessage::Assistant { - content: Some(ChatMessageContent::Text(content)), - reasoning_content: None, - refusal: None, - name: None, - audio: None, - tool_calls: Some(tool_vec), - }, - finish_reason: Some(FinishReason::ToolCalls), - logprobs: None, - } - } else { - ChatCompletionChoice { - index: 0, - message: ChatMessage::Assistant { - content: Some(ChatMessageContent::Text(res)), - reasoning_content: None, - refusal: None, - name: None, - audio: None, - tool_calls: None, - }, - finish_reason: Some(FinishReason::StopSequenceReached), - logprobs: None, - } - }; - response.choices.push(choice); - response -} - -pub fn build_completion_response( - res: String, - model_name: &str, - completion_tokens: Option, - prompt_tokens: Option, -) -> ChatCompletionResponse { - let usage = if prompt_tokens.is_none() && completion_tokens.is_none() { - None - } else { - Some(Usage { - prompt_tokens, - prompt_secs: None, - completion_tokens, - completion_secs: None, - completion_per_token_secs: None, - completion_tps: None, - total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0), - prompt_tokens_details: None, - completion_tokens_details: None, - }) - }; - - build_response(res, model_name, usage) -} - -pub fn build_completion_response_with_time( - res: String, - model_name: &str, - completion_tokens: Option, - completion_secs: Option, - prompt_tokens: Option, - prompt_secs: Option, -) -> ChatCompletionResponse { - let usage = if prompt_tokens.is_none() && completion_tokens.is_none() { - None - } else { - let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) = - completion_tokens - && let Some(completion_secs) = completion_secs - { - let per_token_secs = completion_secs / completion_tokens as f64; - let tps = completion_tokens as f64 / completion_secs; - (Some(per_token_secs), Some(tps)) - } else { - (None, None) - }; - Some(Usage { - prompt_tokens, - prompt_secs, - completion_tokens, - completion_secs, - completion_per_token_secs, - completion_tps, - total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0), - prompt_tokens_details: None, - completion_tokens_details: None, - }) - }; - - build_response(res, model_name, usage) -} - -pub fn build_completion_chunk_response( - res: String, - model_name: &str, - tool_call_id: Option, - tool_call_content: Option, -) -> ChatCompletionChunkResponse { - let id = uuid::Uuid::new_v4().to_string(); - let mut response = ChatCompletionChunkResponse { - id: Some(id), - choices: vec![], - created: timestamp() as u32, - model: model_name.to_string(), - system_fingerprint: None, - object: "chat.completion.chunk".to_string(), - usage: None, - }; - let choice = if let Some(tool_call_id) = tool_call_id { - let function = if let Some(content) = tool_call_content { - match serde_json::from_str::(&content) { - Ok(json_value) => { - let name = json_value - .get("name") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - - let arguments = json_value.get("arguments").map(|v| v.to_string()); - - DeltaFunction { name, arguments } - } - Err(_) => DeltaFunction { - name: None, - arguments: Some(content), - }, - } - } else { - DeltaFunction { - name: None, - arguments: None, - } - }; - ChatCompletionChunkChoice { - index: Some(0), - delta: DeltaChatMessage::Assistant { - content: None, - reasoning_content: None, - refusal: None, - name: None, - tool_calls: Some(vec![DeltaToolCall { - index: Some(0), - id: Some(tool_call_id), - r#type: Some("function".to_string()), - function, - }]), - }, - finish_reason: None, - logprobs: None, - } - } else { - ChatCompletionChunkChoice { - index: Some(0), - delta: DeltaChatMessage::Assistant { - content: Some(ChatMessageContent::Text(res)), - reasoning_content: None, - refusal: None, - name: None, - tool_calls: None, - }, - finish_reason: None, - logprobs: None, - } - }; - response.choices.push(choice); - response -} - pub fn extract_mes(mes: &ChatCompletionParameters) -> Result> { let mut mes_vec = Vec::new(); for chat_mes in mes.messages.clone() { diff --git a/src/utils/response_utils.rs b/src/utils/response_utils.rs new file mode 100644 index 0000000..9ff24a2 --- /dev/null +++ b/src/utils/response_utils.rs @@ -0,0 +1,426 @@ +use crate::{ + params::{ + chat::{ + AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, + ChatCompletionChunkResponse, ChatCompletionResponse, ChatMessage, + ChatMessageAudioContentPart, ChatMessageContent, ChatMessageContentPart, + ChatMessageImageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function, + ImageUrlType, ToolCall, + }, + shared::{FinishReason, Usage}, + }, + utils::timestamp, +}; + +pub fn build_img_completion_response( + base64vec: &Vec, + model_name: &str, +) -> ChatCompletionResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionResponse { + id: Some(id), + choices: vec![], + // created: chrono::Utc::now().timestamp() as u32, + created: timestamp() as u32, + model: model_name.to_string(), + service_tier: None, + system_fingerprint: None, + object: "chat.completion".to_string(), + usage: None, + }; + let mut conten_part_vec = vec![]; + for img_bas64 in base64vec { + let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64; + let part = ChatMessageContentPart::Image(ChatMessageImageContentPart { + r#type: "image".to_string(), + image_url: ImageUrlType { + url: img_base64_prefix, + detail: None, + }, + }); + conten_part_vec.push(part); + } + let choice = ChatCompletionChoice { + index: 0, + message: ChatMessage::Assistant { + content: Some(ChatMessageContent::ContentPart(conten_part_vec)), + reasoning_content: None, + refusal: None, + name: None, + audio: None, + tool_calls: None, + }, + finish_reason: Some(FinishReason::StopSequenceReached), + logprobs: None, + }; + response.choices.push(choice); + response +} + +pub fn build_audio_completion_response( + base64_audio: &String, + model_name: &str, +) -> ChatCompletionResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionResponse { + id: Some(id), + choices: vec![], + created: timestamp() as u32, + model: model_name.to_string(), + service_tier: None, + system_fingerprint: None, + object: "chat.completion".to_string(), + usage: None, + }; + + let base64_audio = format!("data:audio/wav;base64,{}", base64_audio); + let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart { + r#type: "audio".to_string(), + audio_url: AudioUrlType { + url: base64_audio.to_string(), + }, + })]; + let choice = ChatCompletionChoice { + index: 0, + message: ChatMessage::Assistant { + content: Some(ChatMessageContent::ContentPart(conten_part_vec)), + reasoning_content: None, + refusal: None, + name: None, + audio: None, + tool_calls: None, + }, + finish_reason: Some(FinishReason::StopSequenceReached), + logprobs: None, + }; + response.choices.push(choice); + response +} + +/// Builds a chat completion response from the model's output string. +/// +/// This function handles two special formatting patterns in the input: +/// 1. Tool call formatting using the delimiter "" where the content before the first +/// delimiter is treated as the main content, and subsequent parts (separated by "") +/// are parsed as JSON tool call definitions. +/// 2. Reasoning content formatting where content between and tags is +/// extracted as reasoning_content, and the content after becomes the main content. +/// +/// # Arguments +/// * `res` - The raw response string from the model that may contain special formatting +/// * `model_name` - Name of the model generating the response +/// * `usage` - Optional usage statistics to include in the response +/// +/// # Returns +/// A [ChatCompletionResponse](aha/src/params/chat.rs#L11-L32) with processed content, tool calls, and/or reasoning content +/// +/// # Format specifications +/// - Tool call format: Content followed by "" and then JSON-formatted tool call data separated by "" +/// - Reasoning format: Content wrapped in and tags followed by actual response +fn build_response(res: String, model_name: &str, usage: Option) -> ChatCompletionResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionResponse { + id: Some(id), + choices: vec![], + created: timestamp() as u32, + model: model_name.to_string(), + service_tier: None, + system_fingerprint: None, + object: "chat.completion".to_string(), + usage, + }; + let (content, tool_calls) = if res.contains("") { + let mes: Vec<&str> = res.split("").collect(); + let content = mes[0].to_string(); + let mut tool_vec = Vec::new(); + for (i, m) in mes.iter().enumerate().skip(1) { + let tool_mes = m.replace("", ""); + let function = match serde_json::from_str::(&tool_mes) { + Ok(json_value) => { + let name = json_value + .get("name") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()) + .unwrap_or_default(); + + let arguments = json_value + .get("arguments") + .map(|v| v.to_string()) + .unwrap_or_default(); + + Function { name, arguments } + } + Err(_) => Function { + name: "".to_string(), + arguments: "".to_string(), + }, + }; + let tool_call = ToolCall { + id: (i - 1).to_string(), + r#type: "function".to_string(), + function, + }; + tool_vec.push(tool_call); + } + (content, Some(tool_vec)) + } else { + (res, None) + }; + let (content, reasoning_content) = if content.contains("") { + let contents: Vec<&str> = content.split("").collect(); + let reasoning_content = contents[0].to_string().replace("", ""); + let content = contents[1].to_string(); + (content, Some(reasoning_content)) + } else { + (content, None) + }; + let finish_reason = if tool_calls.is_some() { + Some(FinishReason::ToolCalls) + } else { + Some(FinishReason::StopSequenceReached) + }; + let choice = ChatCompletionChoice { + index: 0, + message: ChatMessage::Assistant { + content: Some(ChatMessageContent::Text(content)), + reasoning_content, + refusal: None, + name: None, + audio: None, + tool_calls, + }, + finish_reason, + logprobs: None, + }; + response.choices.push(choice); + response +} + +pub fn build_completion_response( + res: String, + model_name: &str, + completion_tokens: Option, + prompt_tokens: Option, +) -> ChatCompletionResponse { + let usage = if prompt_tokens.is_none() && completion_tokens.is_none() { + None + } else { + Some(Usage { + prompt_tokens, + prompt_secs: None, + completion_tokens, + completion_secs: None, + completion_per_token_secs: None, + completion_tps: None, + total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0), + prompt_tokens_details: None, + completion_tokens_details: None, + }) + }; + + build_response(res, model_name, usage) +} + +pub fn build_completion_response_with_time( + res: String, + model_name: &str, + completion_tokens: Option, + completion_secs: Option, + prompt_tokens: Option, + prompt_secs: Option, +) -> ChatCompletionResponse { + let usage = if prompt_tokens.is_none() && completion_tokens.is_none() { + None + } else { + let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) = + completion_tokens + && let Some(completion_secs) = completion_secs + { + let per_token_secs = completion_secs / completion_tokens as f64; + let tps = completion_tokens as f64 / completion_secs; + (Some(per_token_secs), Some(tps)) + } else { + (None, None) + }; + Some(Usage { + prompt_tokens, + prompt_secs, + completion_tokens, + completion_secs, + completion_per_token_secs, + completion_tps, + total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0), + prompt_tokens_details: None, + completion_tokens_details: None, + }) + }; + + build_response(res, model_name, usage) +} + +pub fn build_chunk_response_with_usage( + model_name: &str, + completion_tokens: Option, + completion_secs: Option, + prompt_tokens: Option, + prompt_secs: Option, +) -> ChatCompletionChunkResponse { + let id = uuid::Uuid::new_v4().to_string(); + let usage = if prompt_tokens.is_none() && completion_tokens.is_none() { + None + } else { + let (completion_per_token_secs, completion_tps) = if let Some(completion_tokens) = + completion_tokens + && let Some(completion_secs) = completion_secs + { + let per_token_secs = completion_secs / completion_tokens as f64; + let tps = if (completion_secs - 0.0).abs() > 0.0001 { + completion_tokens as f64 / completion_secs + } else { + 0.0 + }; + (Some(per_token_secs), Some(tps)) + } else { + (None, None) + }; + Some(Usage { + prompt_tokens, + prompt_secs, + completion_tokens, + completion_secs, + completion_per_token_secs, + completion_tps, + total_tokens: prompt_tokens.unwrap_or(0) + completion_tokens.unwrap_or(0), + prompt_tokens_details: None, + completion_tokens_details: None, + }) + }; + let mut response = ChatCompletionChunkResponse { + id: Some(id), + choices: vec![], + created: timestamp() as u32, + model: model_name.to_string(), + system_fingerprint: None, + object: "chat.completion.chunk".to_string(), + usage, + }; + let choice = ChatCompletionChunkChoice { + index: Some(0), + delta: DeltaChatMessage::Assistant { + content: None, + reasoning_content: None, + refusal: None, + name: None, + tool_calls: None, + }, + finish_reason: None, + logprobs: None, + }; + response.choices.push(choice); + response +} + +pub fn build_chunk_response_with_reasoning( + reasoning: String, + model_name: &str, +) -> ChatCompletionChunkResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionChunkResponse { + id: Some(id), + choices: vec![], + created: timestamp() as u32, + model: model_name.to_string(), + system_fingerprint: None, + object: "chat.completion.chunk".to_string(), + usage: None, + }; + let choice = ChatCompletionChunkChoice { + index: Some(0), + delta: DeltaChatMessage::Assistant { + content: None, + reasoning_content: Some(reasoning), + refusal: None, + name: None, + tool_calls: None, + }, + finish_reason: None, + logprobs: None, + }; + response.choices.push(choice); + response +} + +pub fn build_completion_chunk_response( + res: String, + model_name: &str, + tool_call_id: Option, + tool_call_content: Option, +) -> ChatCompletionChunkResponse { + let id = uuid::Uuid::new_v4().to_string(); + let mut response = ChatCompletionChunkResponse { + id: Some(id), + choices: vec![], + created: timestamp() as u32, + model: model_name.to_string(), + system_fingerprint: None, + object: "chat.completion.chunk".to_string(), + usage: None, + }; + let choice = if let Some(tool_call_id) = tool_call_id { + let function = if let Some(content) = tool_call_content { + match serde_json::from_str::(&content) { + Ok(json_value) => { + let name = json_value + .get("name") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()); + + let arguments = json_value.get("arguments").map(|v| v.to_string()); + + DeltaFunction { name, arguments } + } + Err(_) => DeltaFunction { + name: None, + arguments: Some(content), + }, + } + } else { + DeltaFunction { + name: None, + arguments: None, + } + }; + ChatCompletionChunkChoice { + index: Some(0), + delta: DeltaChatMessage::Assistant { + content: None, + reasoning_content: None, + refusal: None, + name: None, + tool_calls: Some(vec![DeltaToolCall { + index: Some(0), + id: Some(tool_call_id), + r#type: Some("function".to_string()), + function, + }]), + }, + finish_reason: None, + logprobs: None, + } + } else { + ChatCompletionChunkChoice { + index: Some(0), + delta: DeltaChatMessage::Assistant { + content: Some(ChatMessageContent::Text(res)), + reasoning_content: None, + refusal: None, + name: None, + tool_calls: None, + }, + finish_reason: None, + logprobs: None, + } + }; + response.choices.push(choice); + response +} diff --git a/tests/test_deepseek_ocr.rs b/tests/test_deepseek_ocr.rs index 0da515a..492d9c1 100644 --- a/tests/test_deepseek_ocr.rs +++ b/tests/test_deepseek_ocr.rs @@ -89,17 +89,11 @@ fn deepseek_ocr_generate() -> Result<()> { let mut model = DeepseekOCRGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs index c36f222..c17cead 100644 --- a/tests/test_fun_asr_nano.rs +++ b/tests/test_fun_asr_nano.rs @@ -38,17 +38,11 @@ fn fun_asr_nano_generate() -> Result<()> { let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = fun_asr_model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_gelab_zero.rs b/tests/test_gelab_zero.rs index 131e227..6c61a02 100644 --- a/tests/test_gelab_zero.rs +++ b/tests/test_gelab_zero.rs @@ -62,16 +62,10 @@ fn gelab_zero_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = qwen3vl.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_gguf_qwen3_5.rs b/tests/test_gguf_qwen3_5.rs index c8133fc..f17f3a9 100644 --- a/tests/test_gguf_qwen3_5.rs +++ b/tests/test_gguf_qwen3_5.rs @@ -80,16 +80,10 @@ fn gguf_test() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = gguf_qwen3_5.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs index d0df313..0d75e64 100644 --- a/tests/test_glm_asr_nano.rs +++ b/tests/test_glm_asr_nano.rs @@ -39,23 +39,17 @@ fn glm_asr_nano_generate() -> Result<()> { let mut glm_asr_model = GlmAsrNanoGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = glm_asr_model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } #[tokio::test] async fn glm_asr_nano_stream() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_stream -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test test_glm_asr_nano glm_asr_nano_stream -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; let model_path = format!("{}/ZhipuAI/GLM-ASR-Nano-2512/", save_dir); diff --git a/tests/test_glm_ocr.rs b/tests/test_glm_ocr.rs index b9b7a6c..c425ba6 100644 --- a/tests/test_glm_ocr.rs +++ b/tests/test_glm_ocr.rs @@ -39,17 +39,11 @@ fn glm_ocr_generate() -> Result<()> { let mut model = GlmOcrGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_health_models.rs b/tests/test_health_models.rs deleted file mode 100644 index 8c3d569..0000000 --- a/tests/test_health_models.rs +++ /dev/null @@ -1,76 +0,0 @@ -use aha::models::common::model_mapping::WhichModel; - -// Import helper functions from api module - these will need to be made public -// or tested through integration testing - -#[test] -fn test_model_type_classification() { - // Since get_model_type and get_model_id are private to api.rs, - // we document the expected behavior here for reference: - // - // LLM models: MiniCPM4_0_5B, Qwen2_5VL3B, Qwen2_5VL7B, Qwen3_0_6B, - // Qwen3VL2B, Qwen3VL4B, Qwen3VL8B, Qwen3VL32B - // OCR models: DeepSeekOCR, HunyuanOCR, PaddleOCRVL - // ASR models: Qwen3ASR0_6B, Qwen3ASR1_7B, GlmASRNano2512, FunASRNano2512 - // Image models: RMBG2_0, VoxCPM, VoxCPM1_5 - - // This test documents the expected model type classification - let llm_models = [ - WhichModel::MiniCPM4_0_5B, - WhichModel::Qwen2_5VL3B, - WhichModel::Qwen2_5VL7B, - WhichModel::Qwen3_0_6B, - WhichModel::Qwen3VL2B, - WhichModel::Qwen3VL4B, - WhichModel::Qwen3VL8B, - WhichModel::Qwen3VL32B, - ]; - - let ocr_models = [ - WhichModel::DeepSeekOCR, - WhichModel::HunyuanOCR, - WhichModel::PaddleOCRVL, - ]; - - let asr_models = [ - WhichModel::Qwen3ASR0_6B, - WhichModel::Qwen3ASR1_7B, - WhichModel::GlmASRNano2512, - WhichModel::FunASRNano2512, - ]; - - let image_models = [ - WhichModel::RMBG2_0, - WhichModel::VoxCPM, - WhichModel::VoxCPM1_5, - ]; - - // Verify counts - assert_eq!(llm_models.len(), 8); - assert_eq!(ocr_models.len(), 3); - assert_eq!(asr_models.len(), 4); - assert_eq!(image_models.len(), 3); - - // Total models - assert_eq!( - llm_models.len() + ocr_models.len() + asr_models.len() + image_models.len(), - 18 - ); -} - -// Note: Integration tests for the /health and /models endpoints -// should be done with a running server. These would typically: -// -// 1. Start the server with a test model -// 2. Make HTTP requests to /health and /models -// 3. Verify the response format and status codes -// -// Example (pseudo-code): -// -// #[tokio::test] -// async fn test_health_endpoint() { -// let resp = reqwest::get("http://localhost:10100/health").await.unwrap(); -// assert_eq!(resp.status(), 200); -// let json: serde_json::Value = resp.json().await.unwrap(); -// assert_eq!(json["status"], "ok"); -// } diff --git a/tests/test_hunyuan_ocr.rs b/tests/test_hunyuan_ocr.rs index c750cc7..48242e1 100644 --- a/tests/test_hunyuan_ocr.rs +++ b/tests/test_hunyuan_ocr.rs @@ -39,18 +39,11 @@ fn hunyuan_ocr_generate() -> Result<()> { let mut model = HunyuanOCRGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); - Ok(()) } diff --git a/tests/test_lfm2.rs b/tests/test_lfm2.rs index 3c678a8..2659a35 100644 --- a/tests/test_lfm2.rs +++ b/tests/test_lfm2.rs @@ -31,17 +31,11 @@ fn lfm2_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_lfm2vl.rs b/tests/test_lfm2vl.rs index b095471..8ab7150 100644 --- a/tests/test_lfm2vl.rs +++ b/tests/test_lfm2vl.rs @@ -43,18 +43,11 @@ fn lfm2vl_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); - Ok(()) } diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs index 9e3753a..671a383 100644 --- a/tests/test_minicpm4.rs +++ b/tests/test_minicpm4.rs @@ -33,18 +33,11 @@ fn minicpm_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); - Ok(()) } diff --git a/tests/test_paddleocr_vl.rs b/tests/test_paddleocr_vl.rs index 640ad5a..a6ff2e6 100644 --- a/tests/test_paddleocr_vl.rs +++ b/tests/test_paddleocr_vl.rs @@ -89,17 +89,11 @@ fn paddleocr_vl_generate() -> Result<()> { let mut model = PaddleOCRVLGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_qwen2_5vl.rs b/tests/test_qwen2_5vl.rs index b4ae2f6..120edd8 100644 --- a/tests/test_qwen2_5vl.rs +++ b/tests/test_qwen2_5vl.rs @@ -46,18 +46,11 @@ fn qwen2_5vl_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); - Ok(()) } diff --git a/tests/test_qwen3.rs b/tests/test_qwen3.rs index fd94b51..895daf0 100644 --- a/tests/test_qwen3.rs +++ b/tests/test_qwen3.rs @@ -21,7 +21,8 @@ fn qwen3_0_6b_generate() -> Result<()> { "role": "user", "content": "你好啊,你是谁" } - ] + ], + "enable_thinking": true } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; @@ -30,17 +31,11 @@ fn qwen3_0_6b_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } @@ -61,7 +56,8 @@ async fn qwen3_0_6b_stream() -> Result<()> { "role": "user", "content": "你是谁" } - ] + ], + "enable_thinking": true } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; diff --git a/tests/test_qwen3_5.rs b/tests/test_qwen3_5.rs index d901890..176192a 100644 --- a/tests/test_qwen3_5.rs +++ b/tests/test_qwen3_5.rs @@ -33,26 +33,22 @@ fn qwen3_5_generate() -> Result<()> { } ] } - ] + ], + "enable_thinking": true } "#; + // "metadata": {"enable_thinking": "true"} let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = qwen3_5.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } @@ -75,16 +71,17 @@ async fn qwen3_5_stream() -> Result<()> { "type": "image", "image_url": { - "url": "file:///home/jhq/Downloads/gougou1.jpg" + "url": "file://./assets/img/ocr_test3.png" } }, { "type": "text", - "text": "描述这张图片." + "text": "OCR" } ] } - ] + ], + "enable_thinking": true } "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; diff --git a/tests/test_qwen3_asr.rs b/tests/test_qwen3_asr.rs index f477b64..5e750c6 100644 --- a/tests/test_qwen3_asr.rs +++ b/tests/test_qwen3_asr.rs @@ -35,17 +35,11 @@ fn qwen3_asr_generate() -> Result<()> { let mut model = Qwen3AsrGenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = model.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index 039096a..7d0fd61 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -94,17 +94,11 @@ fn qwen3vl_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); let res = qwen3vl.generate(mes)?; - let i_duration = i_start.elapsed(); println!("generate: \n {:?}", res); if let Some(usage) = &res.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); Ok(()) } diff --git a/tests/test_robo_brain.rs b/tests/test_robo_brain.rs index 693e2d2..7da265f 100644 --- a/tests/test_robo_brain.rs +++ b/tests/test_robo_brain.rs @@ -34,17 +34,10 @@ fn robo_brain_generate() -> Result<()> { let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); - let i_start = Instant::now(); - let result = model.generate(mes)?; - let i_duration = i_start.elapsed(); - println!("generate: \n {:?}", result); - if let Some(usage) = &result.usage { - let num_token = usage.total_tokens; - let duration_secs = i_duration.as_secs_f64(); - let tps = num_token as f64 / duration_secs; - println!("Tokens per second (TPS): {:.2}", tps); + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); } - println!("Time elapsed in generate is: {:?}", i_duration); - Ok(()) }