From 35678e3ee4ec1400aa226ecbb86a8f5264d8098f Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Wed, 27 May 2026 00:50:23 +0800 Subject: [PATCH] add MiniCPM5-1B --- README.md | 5 +- README.zh-CN.md | 5 +- docs/changelog.md | 6 ++ docs/changelog.zh-CN.md | 6 ++ docs/supported-models.md | 4 +- docs/supported-models.zh-CN.md | 4 +- src/cli/mod.rs | 3 + src/exec/minicpm4.rs | 4 +- src/exec/minicpm5.rs | 61 +++++++++++++ src/exec/mod.rs | 1 + src/models/common/model_mapping.rs | 3 + src/models/glm_asr_nano/model.rs | 5 +- src/models/llama/mod.rs | 32 ++++++- src/models/minicpm4/generate.rs | 13 ++- src/models/minicpm4/model.rs | 5 +- src/models/minicpm5/config.rs | 26 ++++++ src/models/minicpm5/generate.rs | 135 +++++++++++++++++++++++++++++ src/models/minicpm5/mod.rs | 2 + src/models/mod.rs | 15 +++- tests/config_tests.rs | 11 +++ tests/test_minicpm4.rs | 6 +- tests/test_minicpm5.rs | 41 +++++++++ 22 files changed, 367 insertions(+), 26 deletions(-) create mode 100644 src/exec/minicpm5.rs create mode 100644 src/models/minicpm5/config.rs create mode 100644 src/models/minicpm5/generate.rs create mode 100644 src/models/minicpm5/mod.rs create mode 100644 tests/test_minicpm5.rs diff --git a/README.md b/README.md index f2b261e..9893ee3 100644 --- a/README.md +++ b/README.md @@ -29,7 +29,7 @@ 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, MiniCPM5, 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 | | **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR | @@ -39,6 +39,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an | **Reranker** | Qwen3-Reranker | ## Changelog +### 2026-05-27 +- add MiniCPM5 + ### 2026-05-24 - update doc diff --git a/README.zh-CN.md b/README.zh-CN.md index 17aa3a1..380b5bd 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -28,7 +28,7 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | 类别 | 模型 | |------|------| -| **文本** | Qwen3, MiniCPM4, LFM2, LFM2.5 | +| **文本** | Qwen3, MiniCPM4, MiniCPM5, 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 | | **ASR** | GLM-ASR-Nano, Fun-ASR-Nano, Qwen3-ASR | @@ -38,6 +38,9 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理 | **重排序** | Qwen3-Reranker | ## 更新日志 +### 2026-05-27 +- 新增 MiniCPM5 + ### 2026-05-24 - 更新文档 diff --git a/docs/changelog.md b/docs/changelog.md index 95d6285..9721da2 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -5,6 +5,12 @@ 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-05-27 +- add MiniCPM5 + +### 2026-05-24 +- update doc + ### 2026-05-11 - add Moss-TTS-Nano,its performance is worse than the original Python version diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index ad66452..51f3088 100644 --- a/docs/changelog.zh-CN.md +++ b/docs/changelog.zh-CN.md @@ -5,6 +5,12 @@ 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/), 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 +### 2026-05-27 +- 新增 MiniCPM5 + +### 2026-05-24 +- 更新文档 + ### 2026-05-11 - 添加Moss-TTS-Nano,该模型比较小,计算误差影响较大,效果比python原版差 diff --git a/docs/supported-models.md b/docs/supported-models.md index a8f8d41..e3efc61 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -12,7 +12,8 @@ LiquidAI/LFM2-1.2B LiquidAI llm ✔ LiquidAI/LFM2.5-1.2B-Instruct LiquidAI llm ✔ LiquidAI/LFM2.5-VL-1.6B LiquidAI vlm ✔ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔ -OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔ +OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔ +OpenBMB/MiniCPM5-1B OpenBMB llm ✔ Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔ Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm Qwen/Qwen3-0.6B Qwen llm ✔ @@ -54,6 +55,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔ |-------|-----------|--------|---------| | **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B
Qwen/Qwen3-1.7B
Qwen/Qwen3-4B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **MiniCPM4-0.5B** | 0.5B | OpenBMB/MiniCPM4-0.5B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **MiniCPM5-1B** | 1B | OpenBMB/MiniCPM5-1B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **LFM2-1.2B** | 1.2B | LiquidAI/LFM2-1.2B | [lfm1.0](https://huggingface.co/LiquidAI/LFM2-1.2B/blob/main/LICENSE) | | **LFM2.5-1.2B-Instruct** | 1.2B | LiquidAI/LFM2.5-1.2B-Instruct | [lfm1.0](https://huggingface.co/LiquidAI/LFM2.5-1.2B-Instruct/blob/main/LICENSE) | diff --git a/docs/supported-models.zh-CN.md b/docs/supported-models.zh-CN.md index d6c219c..0395370 100644 --- a/docs/supported-models.zh-CN.md +++ b/docs/supported-models.zh-CN.md @@ -12,7 +12,8 @@ LiquidAI/LFM2-1.2B LiquidAI llm ✔ LiquidAI/LFM2.5-1.2B-Instruct LiquidAI llm ✔ LiquidAI/LFM2.5-VL-1.6B LiquidAI vlm ✔ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔ -OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔ +OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔ +OpenBMB/MiniCPM5-1B OpenBMB llm ✔ Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔ Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm Qwen/Qwen3-0.6B Qwen llm ✔ @@ -54,6 +55,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔ |------|--------|------|---------| | **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B
Qwen/Qwen3-1.7B
Qwen/Qwen3-4B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **MiniCPM4-0.5B** | 0.5B | OpenBMB/MiniCPM4-0.5B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **MiniCPM5-1B** | 1B | OpenBMB/MiniCPM5-1B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **LFM2-1.2B** | 1.2B | LiquidAI/LFM2-1.2B | [lfm1.0](https://huggingface.co/LiquidAI/LFM2-1.2B/blob/main/LICENSE) | | **LFM2.5-1.2B-Instruct** | 1.2B | LiquidAI/LFM2.5-1.2B-Instruct | [lfm1.0](https://huggingface.co/LiquidAI/LFM2.5-1.2B-Instruct/blob/main/LICENSE) | diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 6f71100..d0a7cc9 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -244,6 +244,9 @@ pub(crate) fn run_run(args: RunArgs) -> anyhow::Result<()> { WhichModel::MiniCPM4_0_5B => { minicpm4::MiniCPM4Exec::run(&input, output.as_deref(), &weight_path)?; } + WhichModel::MiniCPM5_1B => { + minicpm5::MiniCPM5Exec::run(&input, output.as_deref(), &weight_path)?; + } WhichModel::LFM2_1_2B | WhichModel::LFM2_5_1_2BInstruct => { lfm2::Lfm2Exec::run(&input, output.as_deref(), &weight_path)?; } diff --git a/src/exec/minicpm4.rs b/src/exec/minicpm4.rs index 5e36a4f..17046d6 100644 --- a/src/exec/minicpm4.rs +++ b/src/exec/minicpm4.rs @@ -5,7 +5,7 @@ use std::time::Instant; use anyhow::{Ok, Result}; use crate::exec::ExecModel; -use crate::models::{GenerateModel, minicpm4::generate::MiniCPMGenerateModel}; +use crate::models::{GenerateModel, minicpm4::generate::MiniCPM4GenerateModel}; use crate::utils::get_file_path; pub struct MiniCPM4Exec; @@ -22,7 +22,7 @@ impl ExecModel for MiniCPM4Exec { }; let i_start = Instant::now(); - let mut model = MiniCPMGenerateModel::init(weight_path, None, None)?; + let mut model = MiniCPM4GenerateModel::init(weight_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); diff --git a/src/exec/minicpm5.rs b/src/exec/minicpm5.rs new file mode 100644 index 0000000..c7247f8 --- /dev/null +++ b/src/exec/minicpm5.rs @@ -0,0 +1,61 @@ +//! MiniCPM4-0.5B exec implementation for CLI `run` subcommand + +use std::time::Instant; + +use anyhow::{Ok, Result}; + +use crate::exec::ExecModel; +use crate::models::GenerateModel; +use crate::models::minicpm5::generate::MiniCPM5GenerateModel; +use crate::utils::get_file_path; + +pub struct MiniCPM5Exec; + +impl ExecModel for MiniCPM5Exec { + fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> { + let input_text = &input[0]; + let target_text = if input_text.starts_with("file://") { + // let path = &input[7..]; + let path = get_file_path(input_text)?; + std::fs::read_to_string(path)? + } else { + input_text.to_string() + }; + + let i_start = Instant::now(); + let mut model = MiniCPM5GenerateModel::init(weight_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let message = format!( + r#"{{ + "temperature": 0.3, + "top_p": 0.8, + "model": "minicpm5", + "messages": [ + {{ + "role": "user", + "content": "{}" + }} + ] + }}"#, + target_text.replace('"', "\\\"") + ); + let mes = serde_json::from_str(&message)?; + + let i_start = Instant::now(); + let result = model.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + // Print result + println!("Result: {:?}", result); + + if let Some(out) = output { + std::fs::write(out, format!("{:?}", result))?; + println!("Output saved to: {}", out); + } + + Ok(()) + } +} diff --git a/src/exec/mod.rs b/src/exec/mod.rs index 40947cc..5f90e44 100644 --- a/src/exec/mod.rs +++ b/src/exec/mod.rs @@ -12,6 +12,7 @@ pub mod hunyuan_ocr; pub mod lfm2; pub mod lfm2vl; pub mod minicpm4; +pub mod minicpm5; pub mod paddleocr_vl; pub mod qwen2_5vl; pub mod qwen3; diff --git a/src/models/common/model_mapping.rs b/src/models/common/model_mapping.rs index dc06db2..009833b 100644 --- a/src/models/common/model_mapping.rs +++ b/src/models/common/model_mapping.rs @@ -16,6 +16,8 @@ pub enum WhichModel { LFM2VL1_6B, #[value(name = "OpenBMB/MiniCPM4-0.5B")] MiniCPM4_0_5B, + #[value(name = "OpenBMB/MiniCPM5-1B")] + MiniCPM5_1B, #[value(name = "Qwen/Qwen2.5-VL-3B-Instruct")] Qwen2_5VL3B, #[value(name = "Qwen/Qwen2.5-VL-7B-Instruct")] @@ -136,6 +138,7 @@ impl WhichModel { match self { // LLM models WhichModel::MiniCPM4_0_5B + | WhichModel::MiniCPM5_1B | WhichModel::Qwen3_0_6B | WhichModel::Qwen3_1_7B | WhichModel::Qwen3_4B diff --git a/src/models/glm_asr_nano/model.rs b/src/models/glm_asr_nano/model.rs index 048baa5..67058b2 100644 --- a/src/models/glm_asr_nano/model.rs +++ b/src/models/glm_asr_nano/model.rs @@ -270,6 +270,7 @@ impl GlmAsrNanoModel { "input_layernorm", "post_attention_layernorm", config.text_config.rope_parameters.rope_theta, + eos_ids.clone(), )?; Ok(Self { config, @@ -317,7 +318,9 @@ impl GlmAsrNanoModel { 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)?; + let logits = self + .language_model + .forward_embeds(&inputs_embeds, seqlen_offset)?; Ok(logits) } pub fn clear_kv_cache(&mut self) { diff --git a/src/models/llama/mod.rs b/src/models/llama/mod.rs index dff27e0..7d8f8c4 100644 --- a/src/models/llama/mod.rs +++ b/src/models/llama/mod.rs @@ -1,5 +1,6 @@ use crate::{ - models::common::modules::NaiveAttnGateUpDownMLPBlock, position_embed::rope::RoPE, + models::common::{InferenceModel, modules::NaiveAttnGateUpDownMLPBlock}, + position_embed::rope::RoPE, utils::tensor_utils::prepare_causal_attention_mask, }; use anyhow::Result; @@ -106,6 +107,7 @@ impl LlamaModel { pub struct LlamaForCausalLM { pub model: LlamaModel, lm_head: Linear, + stop_token_ids: Vec, } impl LlamaForCausalLM { @@ -128,6 +130,7 @@ impl LlamaForCausalLM { input_norm_pp_name: &str, post_norm_pp_name: &str, rope_theta_base: f32, + eos_ids: Vec, ) -> Result { let model = LlamaModel::new( vb.pp("model"), @@ -150,10 +153,18 @@ impl LlamaForCausalLM { rope_theta_base, )?; let lm_head = linear_no_bias(hidden_size, vocab_size, vb.pp("lm_head"))?; - Ok(Self { model, lm_head }) + Ok(Self { + model, + lm_head, + stop_token_ids: eos_ids, + }) } - pub fn forward(&mut self, inputs_embeds: &Tensor, seqlen_offset: usize) -> Result { + pub fn forward_embeds( + &mut self, + inputs_embeds: &Tensor, + seqlen_offset: usize, + ) -> Result { let outputs = self.model.forward(inputs_embeds, seqlen_offset)?; let seq_len = outputs.dim(1)?; let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; @@ -164,3 +175,18 @@ impl LlamaForCausalLM { self.model.clear_kv_cache(); } } + +impl InferenceModel for LlamaForCausalLM { + fn forward_step(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result { + let input_embeds = self.model.embed_tokens.forward(input_ids)?; + self.forward_embeds(&input_embeds, 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/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 98a5758..de3a0ce 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -16,7 +16,7 @@ use crate::models::minicpm4::model::MiniCPMModel; use crate::utils::{find_type_files, get_device, get_dtype}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; -pub struct MiniCPMGenerateModel<'a> { +pub struct MiniCPM4GenerateModel<'a> { chat_template: ChatTemplate<'a>, tokenizer: TokenizerModel, minicpm: MiniCPMModel, @@ -24,7 +24,7 @@ pub struct MiniCPMGenerateModel<'a> { model_name: String, } -impl<'a> MiniCPMGenerateModel<'a> { +impl<'a> MiniCPM4GenerateModel<'a> { pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { let chat_template = ChatTemplate::init(path)?; let tokenizer = TokenizerModel::init(path)?; @@ -33,18 +33,15 @@ impl<'a> MiniCPMGenerateModel<'a> { let device = &get_device(device); let cfg_dtype = cfg.torch_dtype.as_str(); let dtype = get_dtype(dtype, cfg_dtype); - let endoftext_id = cfg.eos_token_id[0]; - 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 eos_ids = vec![endoftext_id, im_end_id]; - let minicpm = MiniCPMModel::new(vb, cfg, eos_ids)?; + let minicpm = MiniCPMModel::new(vb, cfg)?; let model_name = std::path::Path::new(path) .file_name() .and_then(|s| s.to_str()) .unwrap_or("minicpm4") .to_string(); - Ok(MiniCPMGenerateModel { + Ok(MiniCPM4GenerateModel { chat_template, tokenizer, minicpm, @@ -54,7 +51,7 @@ impl<'a> MiniCPMGenerateModel<'a> { } } -impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { +impl<'a> GenerateModel for MiniCPM4GenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { let mes_render = self.chat_template.apply_chat_template(&mes)?; let input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; diff --git a/src/models/minicpm4/model.rs b/src/models/minicpm4/model.rs index 3f8d967..c00343c 100644 --- a/src/models/minicpm4/model.rs +++ b/src/models/minicpm4/model.rs @@ -214,7 +214,7 @@ pub struct MiniCPMModel { } impl MiniCPMModel { - pub fn new(vb: VarBuilder, cfg: MiniCPM4Config, eos_ids: Vec) -> Result { + pub fn new(vb: VarBuilder, cfg: MiniCPM4Config) -> 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); @@ -226,6 +226,7 @@ impl MiniCPMModel { let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?; let rope_emb = MiniCPMLongRoPE::new(&cfg, vb.device())?; let lm_head = Linear::new(embed_tokens.embeddings().clone(), None); + let stop_token_ids = cfg.eos_token_id.clone(); Ok(Self { cfg, embed_tokens, @@ -233,7 +234,7 @@ impl MiniCPMModel { norm, rope_emb, lm_head, - stop_token_ids: eos_ids, + stop_token_ids, }) } diff --git a/src/models/minicpm5/config.rs b/src/models/minicpm5/config.rs new file mode 100644 index 0000000..3469c52 --- /dev/null +++ b/src/models/minicpm5/config.rs @@ -0,0 +1,26 @@ +use candle_nn::Activation; +use serde::Deserialize; + +#[derive(Debug, Clone, PartialEq, Deserialize)] +pub struct MiniCPM5Config { + pub bos_token_id: u32, + pub eos_token_id: Vec, + pub pad_token_id: u32, + pub hidden_act: Activation, + pub hidden_size: usize, + pub initializer_range: f64, + pub intermediate_size: usize, + pub max_position_embeddings: usize, + pub model_type: String, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + pub head_dim: usize, + pub rms_norm_eps: f64, + pub rope_theta: f32, + pub rope_scaling: Option, // Using Value for null/complex scaling objects + pub tie_word_embeddings: bool, + pub torch_dtype: String, + pub use_cache: bool, + pub vocab_size: usize, +} diff --git a/src/models/minicpm5/generate.rs b/src/models/minicpm5/generate.rs new file mode 100644 index 0000000..eaf14e8 --- /dev/null +++ b/src/models/minicpm5/generate.rs @@ -0,0 +1,135 @@ +use crate::models::common::MultiModalData; +use crate::models::common::generate::{ + GenerationContext, generate_generic, generate_stream_generic, +}; +use crate::models::llama::LlamaForCausalLM; +use crate::models::minicpm5::config::MiniCPM5Config; +use crate::params::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; +use anyhow::Result; +use candle_core::{DType, Device}; +use candle_nn::VarBuilder; +use rocket::futures::Stream; + +use crate::utils::{find_type_files, get_device, get_dtype}; +use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; + +pub struct MiniCPM5GenerateModel<'a> { + chat_template: ChatTemplate<'a>, + tokenizer: TokenizerModel, + model: LlamaForCausalLM, + device: Device, + model_name: String, +} + +impl<'a> MiniCPM5GenerateModel<'a> { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let chat_template = ChatTemplate::init(path)?; + let tokenizer = TokenizerModel::init(path)?; + let config_path = path.to_string() + "/config.json"; + let cfg: MiniCPM5Config = serde_json::from_slice(&std::fs::read(config_path)?)?; + let device = &get_device(device); + let cfg_dtype = cfg.torch_dtype.as_str(); + 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 model = LlamaForCausalLM::new( + vb, + cfg.vocab_size, + cfg.hidden_size, + cfg.num_hidden_layers, + cfg.num_attention_heads, + Some(cfg.num_key_value_heads), + Some(cfg.head_dim), + false, + "self_attn", + Some("o_proj"), + cfg.intermediate_size, + cfg.hidden_act, + false, + "mlp", + cfg.rms_norm_eps, + "input_layernorm", + "post_attention_layernorm", + cfg.rope_theta, + cfg.eos_token_id.clone(), + )?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("minicpm5") + .to_string(); + Ok(MiniCPM5GenerateModel { + chat_template, + tokenizer, + model, + device: device.clone(), + model_name, + }) + } +} + +impl<'a> GenerateModel for MiniCPM5GenerateModel<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let mes_render = self.chat_template.apply_chat_template(&mes)?; + 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); + let mut ctx = GenerationContext::new( + mes.temperature, + mes.top_p, + None, + mes.repeat_penalty, + mes.repeat_last_n, + seed, + input_ids.dim(1)?, + sample_len, + self.device.clone(), + ); + + let data = MultiModalData::new(vec![]); + generate_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + &mut ctx, + &self.model_name, + ) + } + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let seed = mes.seed.unwrap_or(34562) as u64; + let mes_render = self.chat_template.apply_chat_template(&mes)?; + 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 = generate_stream_generic( + &mut self.model, + &self.tokenizer, + input_ids, + data, + mes.temperature, + mes.top_p, + None, + mes.repeat_penalty, + mes.repeat_last_n, + seed, + sample_len, + false, + &self.device, + &self.model_name, + )?; + Ok(Box::new(Box::pin(stream))) + } +} diff --git a/src/models/minicpm5/mod.rs b/src/models/minicpm5/mod.rs new file mode 100644 index 0000000..9150fe7 --- /dev/null +++ b/src/models/minicpm5/mod.rs @@ -0,0 +1,2 @@ +pub mod config; +pub mod generate; diff --git a/src/models/mod.rs b/src/models/mod.rs index 1ffaa08..482257a 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -12,6 +12,7 @@ pub mod lfm2; pub mod lfm2vl; pub mod mask_gct; pub mod minicpm4; +pub mod minicpm5; pub mod moss_audio_tokenizer_nano; pub mod moss_tts_nano; pub mod paddleocr_vl; @@ -36,6 +37,7 @@ use crate::{ models::{ all_minilm_l6_v2::AllMiniLML6V2Embedding, common::{embedding::TextEmbedding, model_mapping::WhichModel, reranker::TextRerank}, + minicpm5::generate::MiniCPM5GenerateModel, qwen3_embedding::Qwen3Embedding, qwen3_reranker::Qwen3Reranker, }, @@ -50,7 +52,7 @@ use crate::models::{ fun_asr_nano::generate::FunAsrNanoGenerateModel, glm_asr_nano::generate::GlmAsrNanoGenerateModel, glm_ocr::generate::GlmOcrGenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel, lfm2::generate::Lfm2GenerateModel, - lfm2vl::generate::Lfm2VLGenerateModel, minicpm4::generate::MiniCPMGenerateModel, + lfm2vl::generate::Lfm2VLGenerateModel, minicpm4::generate::MiniCPM4GenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel, qwen3_asr::generate::Qwen3AsrGenerateModel, qwen3vl::generate::Qwen3VLGenerateModel, @@ -74,7 +76,8 @@ pub trait GenerateModel { pub enum ModelInstance<'a> { AllMiniLML6V2(AllMiniLML6V2Embedding), - MiniCPM4(MiniCPMGenerateModel<'a>), + MiniCPM4(MiniCPM4GenerateModel<'a>), + MiniCPM5(MiniCPM5GenerateModel<'a>), Lfm2(Lfm2GenerateModel<'a>), Lfm2VL(Lfm2VLGenerateModel<'a>), Qwen2_5VL(Qwen2_5VLGenerateModel<'a>), @@ -101,6 +104,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { Err(anyhow!("embedding model does not support chat completions")) } ModelInstance::MiniCPM4(model) => model.generate(mes), + ModelInstance::MiniCPM5(model) => model.generate(mes), ModelInstance::Lfm2(model) => model.generate(mes), ModelInstance::Lfm2VL(model) => model.generate(mes), ModelInstance::Qwen2_5VL(model) => model.generate(mes), @@ -143,6 +147,7 @@ impl<'a> GenerateModel for ModelInstance<'a> { Err(anyhow!("embedding model does not support chat completions")) } ModelInstance::MiniCPM4(model) => model.generate_stream(mes), + ModelInstance::MiniCPM5(model) => model.generate_stream(mes), ModelInstance::Lfm2(model) => model.generate_stream(mes), ModelInstance::Lfm2VL(model) => model.generate_stream(mes), ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes), @@ -218,9 +223,13 @@ pub fn load_model<'a>( ModelInstance::AllMiniLML6V2(model) } WhichModel::MiniCPM4_0_5B => { - let model = MiniCPMGenerateModel::init(path, device, dtype)?; + let model = MiniCPM4GenerateModel::init(path, device, dtype)?; ModelInstance::MiniCPM4(model) } + WhichModel::MiniCPM5_1B => { + let model = MiniCPM5GenerateModel::init(path, device, dtype)?; + ModelInstance::MiniCPM5(model) + } WhichModel::LFM2_1_2B | WhichModel::LFM2_5_1_2BInstruct => { let model = Lfm2GenerateModel::init(path, device, dtype)?; ModelInstance::Lfm2(model) diff --git a/tests/config_tests.rs b/tests/config_tests.rs index 85cb815..6c0abb9 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -4,6 +4,7 @@ use aha::models::{ lfm2::config::Lfm2Config, lfm2vl::config::{Lfm2ProcessorConfig, Lfm2VLConfig}, minicpm4::config::MiniCPM4Config, + minicpm5::config::MiniCPM5Config, moss_audio_tokenizer_nano::config::MossAudioTokenizerConfig, moss_tts_nano::config::MossTTSConfig, paddleocr_vl::config::PaddleOCRVLConfig, @@ -139,3 +140,13 @@ fn moss_tts_config() -> Result<()> { println!("{:?}", config); Ok(()) } + +#[test] +fn minicpm5_config() -> Result<()> { + // cargo test -F cuda --test config_tests minicpm5_config -r -- --nocapture + let model_path = "/home/jhq/.aha/OpenBMB/MiniCPM5-1B/"; + let config_path = model_path.to_string() + "/config.json"; + let config: MiniCPM5Config = serde_json::from_slice(&std::fs::read(config_path)?)?; + println!("{:?}", config); + Ok(()) +} diff --git a/tests/test_minicpm4.rs b/tests/test_minicpm4.rs index 671a383..554039c 100644 --- a/tests/test_minicpm4.rs +++ b/tests/test_minicpm4.rs @@ -1,6 +1,6 @@ use std::{pin::pin, time::Instant}; -use aha::models::{GenerateModel, minicpm4::generate::MiniCPMGenerateModel}; +use aha::models::{GenerateModel, minicpm4::generate::MiniCPM4GenerateModel}; use aha::params::chat::ChatCompletionParameters; use anyhow::Result; use rocket::futures::StreamExt; @@ -29,7 +29,7 @@ fn minicpm_generate() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - let mut model = MiniCPMGenerateModel::init(&model_path, None, None)?; + let mut model = MiniCPM4GenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); @@ -62,7 +62,7 @@ async fn minicpm_stream() -> Result<()> { "#; let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); - let mut model = MiniCPMGenerateModel::init(&model_path, None, None)?; + let mut model = MiniCPM4GenerateModel::init(&model_path, None, None)?; let i_duration = i_start.elapsed(); println!("Time elapsed in load model is: {:?}", i_duration); diff --git a/tests/test_minicpm5.rs b/tests/test_minicpm5.rs new file mode 100644 index 0000000..9cea524 --- /dev/null +++ b/tests/test_minicpm5.rs @@ -0,0 +1,41 @@ +use std::time::Instant; + +use aha::{ + models::{GenerateModel, minicpm5::generate::MiniCPM5GenerateModel}, + params::chat::ChatCompletionParameters, +}; +use anyhow::Result; +#[test] +fn minicpm5_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_minicpm5 minicpm5_generate -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/OpenBMB/MiniCPM5-1B/", save_dir); + let message = r#" + { + "temperature": 0.3, + "top_p": 0.8, + "model": "minicpm4", + "messages": [ + { + "role": "user", + "content": "什么是AI" + } + ], + "enable_thinking": true + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut model = MiniCPM5GenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let res = model.generate(mes)?; + println!("generate: \n {:?}", res); + if let Some(usage) = &res.usage { + println!("usage: \n {:?}", usage); + } + Ok(()) +}