add deploy model
This commit is contained in:
+106
@@ -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
@@ -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!");
|
||||
// }
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user