refactor generate code

This commit is contained in:
jhqxxx
2026-04-02 22:22:52 +08:00
parent b254d21efc
commit bd8eee6520
59 changed files with 1745 additions and 1800 deletions
+1 -7
View File
@@ -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(())
}
+1 -7
View File
@@ -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(())
}
+1 -7
View File
@@ -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(())
}
+1 -7
View File
@@ -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(())
}
+2 -8
View File
@@ -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);
+1 -7
View File
@@ -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(())
}
-76
View File
@@ -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");
// }
+1 -8
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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(())
}
+1 -7
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)?;
+1 -7
View File
@@ -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(())
}
+1 -7
View File
@@ -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(())
}
+4 -11
View File
@@ -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(())
}