diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index da698f1..da92d3d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -28,6 +28,11 @@ jobs: - name: Setup toolchain uses: actions-rust-lang/setup-rust-toolchain@v1 + - name: Install make dependencies + run: | + sudo apt-get update + sudo apt-get install -y build-essential + - name: Install FFmpeg dependencies run: | sudo apt-get update @@ -40,8 +45,7 @@ jobs: libavfilter-dev \ libavdevice-dev \ libswresample-dev \ - libswscale-dev - + libswscale-dev - run: make lint @@ -54,6 +58,11 @@ jobs: - name: Setup toolchain uses: actions-rust-lang/setup-rust-toolchain@v1 + - name: Install make dependencies + run: | + sudo apt-get update + sudo apt-get install -y build-essential + - name: Install FFmpeg development packages run: | sudo apt-get update diff --git a/Cargo.lock b/Cargo.lock index b1e249d..e409c12 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -29,10 +29,13 @@ dependencies = [ "candle-nn", "candle-transformers", "chrono", + "clap", + "dirs", "ffmpeg-next", "hound", "image", "minijinja", + "modelscope", "num", "reqwest", "rocket", @@ -103,6 +106,56 @@ dependencies = [ "libc", ] +[[package]] +name = "anstream" +version = "0.6.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "43d5b281e737544384e969a5ccad3f1cdd24b48086a0fc1b2a5262a26b8f4f4a" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5192cca8006f1fd4f7237516f40fa183bb07f8fbdfedaa0036de5ea9b0b45e78" + +[[package]] +name = "anstyle-parse" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e7644824f0aa2c7b9384579234ef10eb7efb6a0deb83f9630a49594dd9c15c2" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e231f6134f61b71076a3eab506c379d4f36122f2af15a9ff04415ea4c3339e2" +dependencies = [ + "windows-sys 0.60.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3e0633414522a32ffaac8ac6cc8f748e090c5717661fddeea04219e2344f5f2a" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.60.2", +] + [[package]] name = "anyhow" version = "1.0.100" @@ -520,12 +573,58 @@ dependencies = [ "libloading", ] +[[package]] +name = "clap" +version = "4.5.51" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4c26d721170e0295f191a69bd9a1f93efcdb0aff38684b61ab5750468972e5f5" +dependencies = [ + "clap_builder", + "clap_derive", +] + +[[package]] +name = "clap_builder" +version = "4.5.51" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75835f0c7bf681bfd05abe44e965760fea999a5286c6eb2d59883634fd02011a" +dependencies = [ + "anstream", + "anstyle", + "clap_lex", + "strsim", +] + +[[package]] +name = "clap_derive" +version = "4.5.49" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a0b5487afeab2deb2ff4e03a807ad1a03ac532ff5a2cee5d86884440c7f7671" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "clap_lex" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1d728cc89cf3aee9ff92b05e62b19ee65a02b5702cff7d5a377e32c6ae29d8d" + [[package]] name = "color_quant" version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d7b894f5411737b7867f4827955924d7c254fc9f4d91a6aad6b097804b1018b" +[[package]] +name = "colorchoice" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b05b61dc5112cbb17e4b6cd61790d9845d13888356391624cbe7e41efeac1e75" + [[package]] name = "compact_str" version = "0.9.0" @@ -554,6 +653,19 @@ dependencies = [ "windows-sys 0.59.0", ] +[[package]] +name = "console" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b430743a6eb14e9764d4260d4c0d8123087d504eeb9c48f2b2a5e810dd369df4" +dependencies = [ + "encode_unicode", + "libc", + "once_cell", + "unicode-width", + "windows-sys 0.61.2", +] + [[package]] name = "cookie" version = "0.18.1" @@ -759,6 +871,27 @@ dependencies = [ "syn", ] +[[package]] +name = "dirs" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3e8aa94d75141228480295a7d0e7feb620b1a5ad9f12bc40be62411e38cce4e" +dependencies = [ + "dirs-sys", +] + +[[package]] +name = "dirs-sys" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e01a3366d27ee9890022452ee61b2b63a67e6f13f58900b651ff5665f0bb1fab" +dependencies = [ + "libc", + "option-ext", + "redox_users", + "windows-sys 0.61.2", +] + [[package]] name = "displaydoc" version = "0.2.5" @@ -1865,13 +1998,26 @@ version = "0.17.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "183b3088984b400f4cfac3620d5e076c84da5364016b4f49473de574b2586235" dependencies = [ - "console", + "console 0.15.11", "number_prefix", "portable-atomic", "unicode-width", "web-time", ] +[[package]] +name = "indicatif" +version = "0.18.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ade6dfcba0dfb62ad59e59e7241ec8912af34fd29e0e743e3db992bd278e8b65" +dependencies = [ + "console 0.16.1", + "portable-atomic", + "unicode-width", + "unit-prefix", + "web-time", +] + [[package]] name = "inlinable_string" version = "0.1.15" @@ -1916,6 +2062,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + [[package]] name = "itertools" version = "0.12.1" @@ -2013,6 +2165,16 @@ version = "0.2.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f9fbbcab51052fe104eb5e5d351cf728d30a5be1fe14d9be8a3b097481fb97de" +[[package]] +name = "libredox" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "416f7e718bdb06000964960ffa43b4335ad4012ae8b99060261aa4a8088d5ccb" +dependencies = [ + "bitflags 2.10.0", + "libc", +] + [[package]] name = "linux-raw-sys" version = "0.11.0" @@ -2167,6 +2329,22 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "modelscope" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fe4bcb2259d8c3b94cb7041a55dcca7a5d16de888b45f7590e389f746e52e560" +dependencies = [ + "anyhow", + "clap", + "futures-util", + "indicatif 0.18.2", + "openssl", + "reqwest", + "serde", + "tokio", +] + [[package]] name = "monostate" version = "0.1.18" @@ -2420,6 +2598,12 @@ version = "1.21.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + [[package]] name = "onig" version = "6.5.1" @@ -2474,6 +2658,15 @@ version = "0.1.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d05e27ee213611ffe7d6348b942e8f942b37114c00cc03cec254295a4a17852e" +[[package]] +name = "openssl-src" +version = "300.5.4+3.5.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a507b3792995dae9b0df8a1c1e3771e8418b7c2d9f0baeba32e6fe8b06c7cb72" +dependencies = [ + "cc", +] + [[package]] name = "openssl-sys" version = "0.9.110" @@ -2482,10 +2675,17 @@ checksum = "0a9f0075ba3c21b09f8e8b2026584b1d18d49388648f2fbbf3c97ea8deced8e2" dependencies = [ "cc", "libc", + "openssl-src", "pkg-config", "vcpkg", ] +[[package]] +name = "option-ext" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d" + [[package]] name = "parking_lot" version = "0.12.5" @@ -2903,6 +3103,17 @@ dependencies = [ "bitflags 2.10.0", ] +[[package]] +name = "redox_users" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4e608c6638b9c18977b00b475ac1f28d14e84b27d8d42f70e0bf1e3dec127ac" +dependencies = [ + "getrandom 0.2.16", + "libredox", + "thiserror 2.0.17", +] + [[package]] name = "ref-cast" version = "1.0.25" @@ -3059,6 +3270,7 @@ dependencies = [ "rocket_codegen", "rocket_http", "serde", + "serde_json", "state", "tempfile", "time", @@ -3699,7 +3911,7 @@ dependencies = [ "derive_builder", "esaxx-rs", "getrandom 0.3.4", - "indicatif", + "indicatif 0.17.11", "itertools 0.14.0", "log", "macro_rules_attribute", @@ -4072,6 +4284,12 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e" +[[package]] +name = "unit-prefix" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "323402cff2dd658f39ca17c789b502021b3f18707c91cdf22e3838e1b4023817" + [[package]] name = "untrusted" version = "0.9.0" @@ -4096,6 +4314,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + [[package]] name = "uuid" version = "1.18.1" diff --git a/Cargo.toml b/Cargo.toml index bf2b905..297cabd 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,9 +24,12 @@ tokenizers = "0.22.1" aha_openai_dive = {version = "1.3.2", features = ["stream"]} uuid = { version = "1.18.1", features = ["v4"]} chrono = "0.4.42" -rocket = "0.5.1" +rocket = { version = "0.5.1", features = ["serde_json", "json"] } tokio = "1.47.1" hound = "3.5.1" +clap = { version = "4.5.51", features = ["derive"] } +modelscope = "0.1.0" +dirs = "6.0.0" [features] flash-attn=["candle-flash-attn"] diff --git a/README.md b/README.md index ea4ae40..aad630b 100644 --- a/README.md +++ b/README.md @@ -44,23 +44,7 @@ aha = { git = "https://github.com/jhqxxx/aha.git", features = ["cuda"] } aha = { git = "https://github.com/jhqxxx/aha.git", features = ["cuda", "flash-attn"] } ``` -### 从源码构建运行测试 -```bash -git clone https://github.com/jhqxxx/aha.git -cd aha -# 修改测试用例中模型路径 -# 运行 Qwen3VL 示例 -cargo test -F cuda qwen3vl_generate -- --nocapture - -# 运行 MiniCPM4 示例 -cargo test -F cuda minicpm_generate -- --nocapture - -# 运行 VoxCPM 示例 -cargo test -F cuda voxcpm_generate -- --nocapture -``` - -## 使用方法 -### VoxCPM示例 +#### VoxCPM使用示例 ```rust use aha::models::voxcpm::generate::VoxCPMGenerate; use aha::utils::audio_utils::save_wav; @@ -88,6 +72,77 @@ fn main() -> Result<()> { } ``` +### 从源码构建运行测试 +```bash +git clone https://github.com/jhqxxx/aha.git +cd aha +# 修改测试用例中模型路径 +# 运行 Qwen3VL 示例 +cargo test -F cuda qwen3vl_generate -- --nocapture + +# 运行 MiniCPM4 示例 +cargo test -F cuda minicpm_generate -- --nocapture + +# 运行 VoxCPM 示例 +cargo test -F cuda voxcpm_generate -- --nocapture +``` + +### 从源码构建部署 +```bash +git clone https://github.com/jhqxxx/aha.git +cd aha +git checkout deploy +``` + +#### cargo run 运行参数说明 +##### 基本用法 +```bash +cargo run -F cuda -- [参数] +``` +##### 参数详解 +1. 端口设置 +----- + -p, --port +* 设置HTTP服务监听的端口号 +* 默认值:10100 +* 示例:--port 8080 或 -p 8080 + +2. 模型选择(必选) +----- + -m, --model +* 指定要加载的模型类型 +* 可选值: + * minicpm4-0.5b:MiniCPM4-0.5B 模型 + * qwen2.5vl-3b:Qwen2.5-VL-3B 模型 + * qwen3vl-2b:Qwen3-VL-2B 模型 +* 示例:--model minicpm4-0.5b 或 -m qwen3vl-2b + +3. 权重路径 +----- + --weight-path +* 指定本地模型权重文件路径 +* 如果指定此参数,则跳过模型下载步骤 +* 示例:--weight-path /path/to/model/dir + +4. 保存路径 +----- + --save-dir +* 指定模型下载保存的目录 +* 默认保存在用户主目录下的 .aha 文件夹中 +* 示例:--save-dir /custom/model/path + +5. 下载重试次数 +----- + --download-retries +* 设置模型下载失败时的最大重试次数 +* 默认值:3次 +* 示例:--download-retries 5 + +##### 注意事项 +* 参数前需要使用双横线 -- 分隔 cargo 命令和应用程序参数 +* 模型参数 (--model 或 -m) 是必需的 +* 如果未指定 --weight-path,程序会自动下载指定模型 +* 下载的模型默认保存在 ~/.aha/ 目录下(除非指定了 --save-dir) ## 开发 ### 项目结构 diff --git a/src/api.rs b/src/api.rs new file mode 100644 index 0000000..38e1359 --- /dev/null +++ b/src/api.rs @@ -0,0 +1,106 @@ +use std::pin::pin; +use std::sync::{Arc, OnceLock}; + +use aha::models::{GenerateModel, ModelInstance, WhichModel, load_model}; +use aha::utils::string_to_static_str; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use rocket::futures::StreamExt; +use rocket::serde::json::Json; +use rocket::{ + Request, + futures::Stream, + http::{ContentType, Status}, + post, + response::{Responder, stream::TextStream}, +}; +use tokio::sync::RwLock; + +static MODEL: OnceLock>>> = OnceLock::new(); + +pub fn init(model_type: WhichModel, path: String) -> anyhow::Result<()> { + let model_path = string_to_static_str(path); + let model = load_model(model_type, model_path)?; + MODEL.get_or_init(|| Arc::new(RwLock::new(model))); + Ok(()) +} + +pub(crate) enum Response + Send> { + Stream(TextStream), + Text(String), + Error(String), +} + +impl<'r, 'o: 'r, R> Responder<'r, 'o> for Response +where + R: Stream + Send + 'o, + 'r: 'o, +{ + fn respond_to(self, req: &'r Request<'_>) -> rocket::response::Result<'o> { + match self { + Response::Stream(stream) => stream.respond_to(req), + Response::Text(text) => text.respond_to(req), + Response::Error(e) => { + let mut res = rocket::response::Response::new(); + res.set_status(Status::InternalServerError); + res.set_header(ContentType::JSON); + res.set_sized_body(e.len(), std::io::Cursor::new(e)); + Ok(res) + } + } + } +} + +#[post("/completions", data = "")] +pub(crate) async fn chat( + req: Json, +) -> (ContentType, Response + Send>) { + match req.stream { + Some(false) => { + let response = { + let model_ref = MODEL + .get() + .cloned() + .ok_or_else(|| anyhow::anyhow!("model not init")) + .unwrap(); + model_ref.write().await.generate(req.into_inner()) + }; + match response { + Ok(res) => { + let response_str = serde_json::to_string(&res).unwrap(); + (ContentType::Text, Response::Text(response_str)) + } + Err(e) => (ContentType::Text, Response::Error(e.to_string())), + } + } + _ => { + let text_stream = TextStream! { + let model_ref = MODEL.get().cloned().ok_or_else(|| anyhow::anyhow!("model not init")).unwrap(); + let mut guard = model_ref.write().await; + let stream_result = guard.generate_stream(req.into_inner()); + match stream_result { + Ok(stream) => { + let mut stream = pin!(stream); + while let Some(result) = stream.next().await { + match result { + Ok(chunk) => { + if let Ok(json_str) = serde_json::to_string(&chunk) { + yield format!("data: {}\n\n", json_str); + } + } + Err(e) => { + yield format!("data: {{\"error\": \"{}\"}}\n\n", e); + break; + } + } + } + yield "data: [DONE]\n\n".to_string(); + }, + Err(e) => { + yield format!("event: error\ndata: {}\n\n", e.to_string()); + } + } + }; + (ContentType::EventStream, Response::Stream(text_stream)) + } + } +} diff --git a/src/main.rs b/src/main.rs index e7a11a9..7e66122 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,3 +1,121 @@ -fn main() { - println!("Hello, world!"); +use std::time::Duration; + +use aha::models::WhichModel; +use clap::Parser; +use dirs::home_dir; +use modelscope::ModelScope; +use rocket::{ + Config, + data::{ByteUnit, Limits}, + routes, +}; +use tokio::time::sleep; + +use crate::api::init; +mod api; + +#[derive(Parser, Debug)] +#[command(version, about, long_about = None)] +struct Args { + #[arg(short, long, default_value_t = 10100)] + port: u16, + + #[arg(short, long)] + model: WhichModel, + + #[arg(long)] + weight_path: Option, + + #[arg(long)] + save_dir: Option, + + #[arg(long)] + download_retries: Option, } +async fn download_model(model_id: &str, save_dir: &str, max_retries: u32) -> anyhow::Result<()> { + let mut attempts = 0u32; + loop { + attempts += 1; + println!( + "Attempting to download model (attempt {}/{})", + attempts, max_retries + ); + + match ModelScope::download(model_id, save_dir).await { + Ok(()) => { + println!("Model downloaded successfully"); + return Ok(()); + } + Err(e) => { + if attempts >= max_retries { + return Err(anyhow::anyhow!( + "Failed to download model after {} attempts. Last error: {}", + max_retries, + e + )); + } + + println!( + "Download failed (attempt {}): {}. Retrying in 2 seconds...", + attempts, e + ); + sleep(Duration::from_secs(2)).await; + } + } + } +} + +fn get_default_save_dir() -> Option { + home_dir().map(|mut path| { + path.push(".aha"); // 在 home 目录下创建 .aha 文件夹 + path.to_string_lossy().to_string() + }) +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + let args = Args::parse(); + let model_id = match args.model { + WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B", + WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct", + WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct", + }; + let model_path = match args.weight_path { + Some(path) => path, + None => { + let save_dir = match args.save_dir { + Some(dir) => dir, + None => get_default_save_dir().expect("Failed to get home directory"), + }; + let max_retries = args.download_retries.unwrap_or(3); + download_model(model_id, &save_dir, max_retries).await?; + save_dir + "/" + model_id + } + }; + println!("-------------------download path: {}", model_path); + init(args.model, model_path)?; + start_http_server(args.port).await?; + + Ok(()) +} + +pub async fn start_http_server(port: u16) -> anyhow::Result<()> { + let mut builder = rocket::build().configure(Config { + port, + limits: Limits::default() + .limit("string", ByteUnit::Mebibyte(5)) + .limit("json", ByteUnit::Mebibyte(5)) + .limit("data-form", ByteUnit::Mebibyte(100)) + .limit("file", ByteUnit::Mebibyte(100)), + ..Config::default() + }); + + builder = builder.mount("/chat", routes![api::chat]); + + builder.launch().await?; + Ok(()) +} + +// fn main() { +// println!("Hello, world!"); +// } diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 6525ea8..c5edabd 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -53,7 +53,11 @@ impl<'a> MiniCPMGenerateModel<'a> { impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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)?; @@ -80,8 +84,19 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { fn generate_stream( &mut self, mes: ChatCompletionParameters, - ) -> Result>> { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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)?; @@ -125,6 +140,6 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { } self.minicpm.clear_kv_cache(); }; - Ok(stream) + Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/mod.rs b/src/models/mod.rs index 5206cac..4bdc359 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -10,12 +10,84 @@ use aha_openai_dive::v1::resources::chat::{ use anyhow::Result; use rocket::futures::Stream; +use crate::models::{ + minicpm4::generate::MiniCPMGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, + qwen3vl::generate::Qwen3VLGenerateModel, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)] +pub enum WhichModel { + #[value(name = "minicpm4-0.5b")] + MiniCPM4_0_5B, + #[value(name = "qwen2.5vl-3b")] + Qwen2_5vl3B, + #[value(name = "qwen3vl-2b")] + Qwen3vl2B, +} + pub trait GenerateModel { fn generate(&mut self, mes: ChatCompletionParameters) -> Result; fn generate_stream( &mut self, mes: ChatCompletionParameters, - ) -> Result>> - where - Self: Sized; + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + >; +} + +pub enum ModelInstance<'a> { + MiniCPM4(MiniCPMGenerateModel<'a>), + Qwen2_5VL(Qwen2_5VLGenerateModel<'a>), + Qwen3VL(Qwen3VLGenerateModel<'a>), +} + +impl<'a> GenerateModel for ModelInstance<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + match self { + ModelInstance::MiniCPM4(model) => model.generate(mes), + ModelInstance::Qwen2_5VL(model) => model.generate(mes), + ModelInstance::Qwen3VL(model) => model.generate(mes), + } + } + + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + match self { + ModelInstance::MiniCPM4(model) => model.generate_stream(mes), + ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes), + ModelInstance::Qwen3VL(model) => model.generate_stream(mes), + } + } +} + +pub fn load_model(model_type: WhichModel, path: &str) -> Result> { + let model = match model_type { + WhichModel::MiniCPM4_0_5B => { + let model = MiniCPMGenerateModel::init(path, None, None)?; + ModelInstance::MiniCPM4(model) + } + WhichModel::Qwen2_5vl3B => { + let model = Qwen2_5VLGenerateModel::init(path, None, None)?; + ModelInstance::Qwen2_5VL(model) + } + WhichModel::Qwen3vl2B => { + let model = Qwen3VLGenerateModel::init(path, None, None)?; + ModelInstance::Qwen3VL(model) + } + }; + Ok(model) } diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 6825a94..7468abe 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -63,7 +63,11 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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 mut input_ids = self @@ -122,8 +126,19 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { fn generate_stream( &mut self, mes: ChatCompletionParameters, - ) -> Result>> { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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 mut input_ids = self @@ -203,6 +218,6 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { } self.qwen2_5_vl.clear_kv_cache(); }; - Ok(stream) + Ok(Box::new(Box::pin(stream))) } } diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 95ff996..632e471 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -76,7 +76,12 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { Some(top_p) => top_p, }; let top_k = self.generation_config.top_k; - let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k)); + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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 input = self.pre_processor.process_info(&mes, &mes_render)?; let mut input_ids = self @@ -123,7 +128,14 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { fn generate_stream( &mut self, mes: ChatCompletionParameters, - ) -> Result>> { + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { let temperature = match mes.temperature { None => self.generation_config.temperature, Some(tem) => tem, @@ -133,7 +145,12 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { Some(top_p) => top_p, }; let top_k = self.generation_config.top_k; - let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k)); + let seed = match mes.seed { + None => 34562u64, + Some(s) => s 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 input = self.pre_processor.process_info(&mes, &mes_render)?; let mut input_ids = self @@ -199,6 +216,6 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { } self.qwen3_vl.clear_kv_cache(); }; - Ok(stream) + Ok(Box::new(Box::pin(stream))) } } diff --git a/src/utils/mod.rs b/src/utils/mod.rs index deb20c0..550e54e 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -255,10 +255,11 @@ pub fn get_logit_processor( temperature: Option, top_p: Option, top_k: Option, + seed: u64, ) -> LogitsProcessor { match top_k { None => LogitsProcessor::new( - 34562, + seed, temperature.map(|temp| temp as f64), top_p.map(|tp| tp as f64), ), @@ -277,7 +278,7 @@ pub fn get_logit_processor( }, }, }; - LogitsProcessor::from_sampling(34562, sampling) + LogitsProcessor::from_sampling(seed, sampling) } } }