refactor generate code
This commit is contained in:
@@ -89,17 +89,11 @@ fn deepseek_ocr_generate() -> Result<()> {
|
||||
let mut model = DeepseekOCRGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -38,17 +38,11 @@ fn fun_asr_nano_generate() -> Result<()> {
|
||||
let mut fun_asr_model = FunAsrNanoGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = fun_asr_model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -62,16 +62,10 @@ fn gelab_zero_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = qwen3vl.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -80,16 +80,10 @@ fn gguf_test() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = gguf_qwen3_5.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -39,23 +39,17 @@ fn glm_asr_nano_generate() -> Result<()> {
|
||||
let mut glm_asr_model = GlmAsrNanoGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = glm_asr_model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn glm_asr_nano_stream() -> Result<()> {
|
||||
// RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_stream -r -- --nocapture
|
||||
// RUST_BACKTRACE=1 cargo test -F cuda --test test_glm_asr_nano glm_asr_nano_stream -r -- --nocapture
|
||||
let save_dir =
|
||||
aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?;
|
||||
let model_path = format!("{}/ZhipuAI/GLM-ASR-Nano-2512/", save_dir);
|
||||
|
||||
@@ -39,17 +39,11 @@ fn glm_ocr_generate() -> Result<()> {
|
||||
let mut model = GlmOcrGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
use aha::models::common::model_mapping::WhichModel;
|
||||
|
||||
// Import helper functions from api module - these will need to be made public
|
||||
// or tested through integration testing
|
||||
|
||||
#[test]
|
||||
fn test_model_type_classification() {
|
||||
// Since get_model_type and get_model_id are private to api.rs,
|
||||
// we document the expected behavior here for reference:
|
||||
//
|
||||
// LLM models: MiniCPM4_0_5B, Qwen2_5VL3B, Qwen2_5VL7B, Qwen3_0_6B,
|
||||
// Qwen3VL2B, Qwen3VL4B, Qwen3VL8B, Qwen3VL32B
|
||||
// OCR models: DeepSeekOCR, HunyuanOCR, PaddleOCRVL
|
||||
// ASR models: Qwen3ASR0_6B, Qwen3ASR1_7B, GlmASRNano2512, FunASRNano2512
|
||||
// Image models: RMBG2_0, VoxCPM, VoxCPM1_5
|
||||
|
||||
// This test documents the expected model type classification
|
||||
let llm_models = [
|
||||
WhichModel::MiniCPM4_0_5B,
|
||||
WhichModel::Qwen2_5VL3B,
|
||||
WhichModel::Qwen2_5VL7B,
|
||||
WhichModel::Qwen3_0_6B,
|
||||
WhichModel::Qwen3VL2B,
|
||||
WhichModel::Qwen3VL4B,
|
||||
WhichModel::Qwen3VL8B,
|
||||
WhichModel::Qwen3VL32B,
|
||||
];
|
||||
|
||||
let ocr_models = [
|
||||
WhichModel::DeepSeekOCR,
|
||||
WhichModel::HunyuanOCR,
|
||||
WhichModel::PaddleOCRVL,
|
||||
];
|
||||
|
||||
let asr_models = [
|
||||
WhichModel::Qwen3ASR0_6B,
|
||||
WhichModel::Qwen3ASR1_7B,
|
||||
WhichModel::GlmASRNano2512,
|
||||
WhichModel::FunASRNano2512,
|
||||
];
|
||||
|
||||
let image_models = [
|
||||
WhichModel::RMBG2_0,
|
||||
WhichModel::VoxCPM,
|
||||
WhichModel::VoxCPM1_5,
|
||||
];
|
||||
|
||||
// Verify counts
|
||||
assert_eq!(llm_models.len(), 8);
|
||||
assert_eq!(ocr_models.len(), 3);
|
||||
assert_eq!(asr_models.len(), 4);
|
||||
assert_eq!(image_models.len(), 3);
|
||||
|
||||
// Total models
|
||||
assert_eq!(
|
||||
llm_models.len() + ocr_models.len() + asr_models.len() + image_models.len(),
|
||||
18
|
||||
);
|
||||
}
|
||||
|
||||
// Note: Integration tests for the /health and /models endpoints
|
||||
// should be done with a running server. These would typically:
|
||||
//
|
||||
// 1. Start the server with a test model
|
||||
// 2. Make HTTP requests to /health and /models
|
||||
// 3. Verify the response format and status codes
|
||||
//
|
||||
// Example (pseudo-code):
|
||||
//
|
||||
// #[tokio::test]
|
||||
// async fn test_health_endpoint() {
|
||||
// let resp = reqwest::get("http://localhost:10100/health").await.unwrap();
|
||||
// assert_eq!(resp.status(), 200);
|
||||
// let json: serde_json::Value = resp.json().await.unwrap();
|
||||
// assert_eq!(json["status"], "ok");
|
||||
// }
|
||||
@@ -39,18 +39,11 @@ fn hunyuan_ocr_generate() -> Result<()> {
|
||||
let mut model = HunyuanOCRGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+4
-10
@@ -31,17 +31,11 @@ fn lfm2_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+4
-11
@@ -43,18 +43,11 @@ fn lfm2vl_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+4
-11
@@ -33,18 +33,11 @@ fn minicpm_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -89,17 +89,11 @@ fn paddleocr_vl_generate() -> Result<()> {
|
||||
let mut model = PaddleOCRVLGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+4
-11
@@ -46,18 +46,11 @@ fn qwen2_5vl_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+8
-12
@@ -21,7 +21,8 @@ fn qwen3_0_6b_generate() -> Result<()> {
|
||||
"role": "user",
|
||||
"content": "你好啊,你是谁"
|
||||
}
|
||||
]
|
||||
],
|
||||
"enable_thinking": true
|
||||
}
|
||||
"#;
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
@@ -30,17 +31,11 @@ fn qwen3_0_6b_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -61,7 +56,8 @@ async fn qwen3_0_6b_stream() -> Result<()> {
|
||||
"role": "user",
|
||||
"content": "你是谁"
|
||||
}
|
||||
]
|
||||
],
|
||||
"enable_thinking": true
|
||||
}
|
||||
"#;
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
|
||||
+8
-11
@@ -33,26 +33,22 @@ fn qwen3_5_generate() -> Result<()> {
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
],
|
||||
"enable_thinking": true
|
||||
}
|
||||
"#;
|
||||
// "metadata": {"enable_thinking": "true"}
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
let i_start = Instant::now();
|
||||
let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = qwen3_5.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -75,16 +71,17 @@ async fn qwen3_5_stream() -> Result<()> {
|
||||
"type": "image",
|
||||
"image_url":
|
||||
{
|
||||
"url": "file:///home/jhq/Downloads/gougou1.jpg"
|
||||
"url": "file://./assets/img/ocr_test3.png"
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "描述这张图片."
|
||||
"text": "OCR"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
],
|
||||
"enable_thinking": true
|
||||
}
|
||||
"#;
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
|
||||
@@ -35,17 +35,11 @@ fn qwen3_asr_generate() -> Result<()> {
|
||||
let mut model = Qwen3AsrGenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
let i_start = Instant::now();
|
||||
let res = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -94,17 +94,11 @@ fn qwen3vl_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = qwen3vl.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
@@ -34,17 +34,10 @@ fn robo_brain_generate() -> Result<()> {
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", result);
|
||||
if let Some(usage) = &result.usage {
|
||||
let num_token = usage.total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
let res = model.generate(mes)?;
|
||||
println!("generate: \n {:?}", res);
|
||||
if let Some(usage) = &res.usage {
|
||||
println!("usage: \n {:?}", usage);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user