add deploy model

This commit is contained in:
jhqxxx
2025-11-05 14:46:03 +08:00
parent a5721e7f50
commit d15cd45315
11 changed files with 676 additions and 41 deletions
+10 -1
View File
@@ -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
@@ -42,7 +47,6 @@ jobs:
libswresample-dev \
libswscale-dev
- run: make lint
build-and-test:
@@ -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
Generated
+226 -2
View File
@@ -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"
+4 -1
View File
@@ -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"]
+72 -17
View File
@@ -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 <PORT>
* 设置HTTP服务监听的端口号
* 默认值:10100
* 示例:--port 8080 或 -p 8080
2. 模型选择(必选)
-----
-m, --model <MODEL>
* 指定要加载的模型类型
* 可选值:
* minicpm4-0.5bMiniCPM4-0.5B 模型
* qwen2.5vl-3bQwen2.5-VL-3B 模型
* qwen3vl-2bQwen3-VL-2B 模型
* 示例:--model minicpm4-0.5b 或 -m qwen3vl-2b
3. 权重路径
-----
--weight-path <WEIGHT_PATH>
* 指定本地模型权重文件路径
* 如果指定此参数,则跳过模型下载步骤
* 示例:--weight-path /path/to/model/dir
4. 保存路径
-----
--save-dir <SAVE_DIR>
* 指定模型下载保存的目录
* 默认保存在用户主目录下的 .aha 文件夹中
* 示例:--save-dir /custom/model/path
5. 下载重试次数
-----
--download-retries <DOWNLOAD_RETRIES>
* 设置模型下载失败时的最大重试次数
* 默认值:3次
* 示例:--download-retries 5
##### 注意事项
* 参数前需要使用双横线 -- 分隔 cargo 命令和应用程序参数
* 模型参数 (--model 或 -m) 是必需的
* 如果未指定 --weight-path,程序会自动下载指定模型
* 下载的模型默认保存在 ~/.aha/ 目录下(除非指定了 --save-dir
## 开发
### 项目结构
+106
View File
@@ -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<Arc<RwLock<ModelInstance<'static>>>> = 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<R: Stream<Item = String> + Send> {
Stream(TextStream<R>),
Text(String),
Error(String),
}
impl<'r, 'o: 'r, R> Responder<'r, 'o> for Response<R>
where
R: Stream<Item = String> + 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 = "<req>")]
pub(crate) async fn chat(
req: Json<ChatCompletionParameters>,
) -> (ContentType, Response<impl Stream<Item = String> + 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))
}
}
}
+120 -2
View File
@@ -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<String>,
#[arg(long)]
save_dir: Option<String>,
#[arg(long)]
download_retries: Option<u32>,
}
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<String> {
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!");
// }
+19 -4
View File
@@ -53,7 +53,11 @@ impl<'a> MiniCPMGenerateModel<'a> {
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
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<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
}
}
+75 -3
View File
@@ -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<ChatCompletionResponse>;
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>>
where
Self: Sized;
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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<ChatCompletionResponse> {
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<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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<ModelInstance<'_>> {
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)
}
+19 -4
View File
@@ -63,7 +63,11 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
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<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
}
}
+21 -4
View File
@@ -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<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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)))
}
}
+3 -2
View File
@@ -255,10 +255,11 @@ pub fn get_logit_processor(
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
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)
}
}
}