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
+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)
}
}
}