主要改了 vendor/esaxx-rs/build.rs:只在 crt-static 目标下才启用 static_crt(true),用来修复 Windows 下常见的 MSVC 运行库冲突
This commit is contained in:
@@ -0,0 +1,69 @@
|
||||
use aha::models::{ArtifactKind, LoadSpec, ModelPaths, WhichModel};
|
||||
|
||||
#[test]
|
||||
fn load_spec_auto_resolves_to_safetensors_default() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Auto,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some("D:/model_download/Qwen3-Embedding-0.6B".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
assert_eq!(spec.resolved_artifact(), ArtifactKind::Safetensors);
|
||||
assert!(spec.validate().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_gguf_requires_gguf_path() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5Gguf,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths::default(),
|
||||
};
|
||||
|
||||
let err = spec.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("gguf_path is required"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_onnx_requires_onnx_path() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths::default(),
|
||||
};
|
||||
|
||||
let err = spec.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("onnx_path is required"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_rejects_unsupported_artifact() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::MiniCPM4_0_5B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some("D:/tmp/model.onnx".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let err = spec.validate().unwrap_err().to_string();
|
||||
assert!(err.contains("does not support artifact"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_accepts_qwen3_5_onnx() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some("D:/model_download/Qwen3.5-0.8B-ONNX".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
assert!(spec.validate().is_ok());
|
||||
}
|
||||
@@ -0,0 +1,301 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use aha::models::{
|
||||
ArtifactKind, GenerateModel, LoadSpec, ModelPaths, WhichModel,
|
||||
qwen3_5::generate::Qwen3_5GenerateModel,
|
||||
};
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
|
||||
const DEFAULT_QWEN3_5_SAFETENSORS_DIR: &str = r"D:\model_download\Qwen3.5-0.8B";
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
const DEFAULT_QWEN3_5_ONNX_DIR: &str = r"D:\model_download\Qwen3.5-0.8B-ONNX";
|
||||
const DEFAULT_QWEN3_5_GGUF_DIRS: &[&str] = &[
|
||||
r"D:\model_download\Qwen3.5-0.8B-GGUF",
|
||||
r"D:\model_download\Qwen3.5-0.8B-gguf",
|
||||
r"D:\model_download\Qwen3.5-2B-GGUF",
|
||||
r"D:\model_download\Qwen3.5-4B-GGUF",
|
||||
];
|
||||
|
||||
fn env_or_default(key: &str, default: &str) -> String {
|
||||
std::env::var(key).unwrap_or_else(|_| default.to_string())
|
||||
}
|
||||
|
||||
fn existing_dir(path: &str) -> bool {
|
||||
let p = Path::new(path);
|
||||
p.exists() && p.is_dir()
|
||||
}
|
||||
|
||||
fn first_file_with_extension_recursive(dir: &str, extension: &str) -> Result<Option<PathBuf>> {
|
||||
if !existing_dir(dir) {
|
||||
return Ok(None);
|
||||
}
|
||||
let mut stack = vec![PathBuf::from(dir)];
|
||||
let mut matches = Vec::new();
|
||||
while let Some(current) = stack.pop() {
|
||||
for entry in std::fs::read_dir(¤t)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if path.is_dir() {
|
||||
stack.push(path);
|
||||
} else if path
|
||||
.extension()
|
||||
.is_some_and(|ext| ext.eq_ignore_ascii_case(extension))
|
||||
{
|
||||
matches.push(path);
|
||||
}
|
||||
}
|
||||
}
|
||||
matches.sort();
|
||||
Ok(matches.into_iter().next())
|
||||
}
|
||||
|
||||
fn resolve_gguf_path() -> Result<Option<String>> {
|
||||
if let Ok(path) = std::env::var("AHA_QWEN3_5_GGUF_PATH")
|
||||
&& Path::new(&path).exists()
|
||||
{
|
||||
return Ok(Some(path));
|
||||
}
|
||||
for dir in DEFAULT_QWEN3_5_GGUF_DIRS {
|
||||
if let Some(path) = first_file_with_extension_recursive(dir, "gguf")? {
|
||||
return Ok(Some(path.to_string_lossy().to_string()));
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn build_text_request() -> Result<ChatCompletionParameters> {
|
||||
let payload = serde_json::json!({
|
||||
"model": "qwen3.5-0.8b",
|
||||
"max_tokens": 8,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "请用一句话介绍 Rust。"
|
||||
}
|
||||
]
|
||||
});
|
||||
Ok(serde_json::from_value(payload)?)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_5_safetensors_init_from_spec_can_generate() -> Result<()> {
|
||||
let weight_dir = env_or_default(
|
||||
"AHA_QWEN3_5_SAFETENSORS_DIR",
|
||||
DEFAULT_QWEN3_5_SAFETENSORS_DIR,
|
||||
);
|
||||
if !existing_dir(&weight_dir) {
|
||||
println!("skip safetensors test: dir not found, set AHA_QWEN3_5_SAFETENSORS_DIR to run");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Safetensors,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some(weight_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3_5GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_5_gguf_init_from_spec_can_generate() -> Result<()> {
|
||||
let Some(gguf_path) = resolve_gguf_path()? else {
|
||||
println!("skip gguf test: no gguf file found, set AHA_QWEN3_5_GGUF_PATH to run explicitly");
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(gguf_path),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3_5GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[test]
|
||||
fn qwen3_5_onnx_init_from_spec_can_generate() -> Result<()> {
|
||||
use aha::models::common::onnx::ensure_ort_dylib_path;
|
||||
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip onnx test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let onnx_dir = env_or_default("AHA_QWEN3_5_ONNX_DIR", DEFAULT_QWEN3_5_ONNX_DIR);
|
||||
if !existing_dir(&onnx_dir) {
|
||||
println!("skip onnx test: dir not found, set AHA_QWEN3_5_ONNX_DIR to run");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(onnx_dir.clone()),
|
||||
tokenizer_dir: Some(onnx_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3_5GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[test]
|
||||
fn qwen3_5_onnx_image_multimodal_can_generate() -> Result<()> {
|
||||
use aha::models::common::onnx::ensure_ort_dylib_path;
|
||||
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip onnx multimodal image test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let onnx_dir = env_or_default("AHA_QWEN3_5_ONNX_DIR", DEFAULT_QWEN3_5_ONNX_DIR);
|
||||
if !existing_dir(&onnx_dir) {
|
||||
println!("skip onnx multimodal image test: dir not found");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(onnx_dir.clone()),
|
||||
tokenizer_dir: Some(onnx_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3_5GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let image_path = std::env::current_dir()?
|
||||
.join("assets")
|
||||
.join("img")
|
||||
.join("ocr_test1.png");
|
||||
if !image_path.exists() {
|
||||
println!(
|
||||
"skip onnx multimodal image test: local image not found at {}",
|
||||
image_path.display()
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
let image_url = format!(
|
||||
"file:///{}",
|
||||
image_path.to_string_lossy().replace('\\', "/")
|
||||
);
|
||||
let payload = serde_json::json!({
|
||||
"model": "qwen3.5-0.8b",
|
||||
"max_tokens": 8,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image",
|
||||
"image_url": {"url": image_url}
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "识别图像中的文字"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
});
|
||||
let request: ChatCompletionParameters = serde_json::from_value(payload)?;
|
||||
let response = model.generate(request)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(
|
||||
choices_len > 0,
|
||||
"expected at least one generated choice for multimodal image request"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[test]
|
||||
fn qwen3_5_onnx_video_multimodal_is_rejected() -> Result<()> {
|
||||
use aha::models::common::onnx::ensure_ort_dylib_path;
|
||||
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip onnx multimodal video rejection test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let onnx_dir = env_or_default("AHA_QWEN3_5_ONNX_DIR", DEFAULT_QWEN3_5_ONNX_DIR);
|
||||
if !existing_dir(&onnx_dir) {
|
||||
println!("skip onnx multimodal video rejection test: dir not found");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_5_0_8B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(onnx_dir.clone()),
|
||||
tokenizer_dir: Some(onnx_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3_5GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let payload = serde_json::json!({
|
||||
"model": "qwen3.5-0.8b",
|
||||
"max_tokens": 8,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "video",
|
||||
"video_url": {"url": "file://./assets/video/dummy.mp4"}
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "描述视频内容"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
});
|
||||
let request: ChatCompletionParameters = serde_json::from_value(payload)?;
|
||||
let err = model
|
||||
.generate(request)
|
||||
.expect_err("onnx backend should reject video multimodal input for now");
|
||||
assert!(
|
||||
err.to_string().contains("audio/video") || err.to_string().contains("video"),
|
||||
"unexpected error: {err}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
@@ -5,11 +5,13 @@ use std::{
|
||||
};
|
||||
|
||||
use aha::models::{
|
||||
common::{gguf::Gguf, retrieval::cosine_similarity},
|
||||
ArtifactKind, LoadSpec, ModelPaths,
|
||||
common::{gguf::Gguf, onnx::ensure_ort_dylib_path, retrieval::cosine_similarity},
|
||||
qwen3_embedding::generate::Qwen3EmbeddingModel,
|
||||
};
|
||||
use anyhow::{Context, Result, anyhow};
|
||||
use candle_core::{Device, quantized::gguf_file};
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use ort::session::Session;
|
||||
|
||||
const QWEN3_EMBEDDING_SAFETENSORS_DIR: &str = r"D:\model_download\Qwen3-Embedding-0.6B";
|
||||
@@ -130,8 +132,7 @@ fn qwen3_embedding_onnx_can_load() -> Result<()> {
|
||||
// Run this test only:
|
||||
// cargo test --test test_qwen3_embedding_multi_format qwen3_embedding_onnx_can_load -- --nocapture
|
||||
let onnx_path = first_file_with_extension_recursive(QWEN3_EMBEDDING_ONNX_DIR, "onnx")?;
|
||||
// Current aha runtime does not integrate ONNX execution yet.
|
||||
// Here we validate that ONNX artifact can be discovered and read normally.
|
||||
// Basic artifact-level smoke check: ONNX file can be discovered and read.
|
||||
let metadata = std::fs::metadata(&onnx_path)
|
||||
.with_context(|| format!("failed to read onnx metadata: {}", onnx_path.display()))?;
|
||||
if metadata.len() == 0 {
|
||||
@@ -143,6 +144,7 @@ fn qwen3_embedding_onnx_can_load() -> Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[test]
|
||||
fn qwen3_embedding_onnxruntime_can_create_session() -> Result<()> {
|
||||
// Run this test only:
|
||||
@@ -152,8 +154,8 @@ fn qwen3_embedding_onnxruntime_can_create_session() -> Result<()> {
|
||||
// ort with `load-dynamic` requires ONNX Runtime dynamic library path to be configured.
|
||||
// Example on Windows:
|
||||
// $env:ORT_DYLIB_PATH = "D:\\onnxruntime\\onnxruntime.dll"
|
||||
if std::env::var("ORT_DYLIB_PATH").is_err() {
|
||||
println!("skip onnxruntime session test: ORT_DYLIB_PATH is not set");
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip onnxruntime session test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
@@ -177,6 +179,92 @@ fn qwen3_embedding_onnxruntime_can_create_session() -> Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_embedding_onnx_init_from_spec_can_embed() -> Result<()> {
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip onnx init test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: aha::models::WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(QWEN3_EMBEDDING_ONNX_DIR.to_string()),
|
||||
tokenizer_dir: Some(QWEN3_EMBEDDING_ONNX_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3EmbeddingModel::init_from_spec(&spec, None, None)?;
|
||||
let output = model.embed(&["test onnx embedding".to_string()])?;
|
||||
assert_eq!(output.len(), 1);
|
||||
assert!(!output[0].is_empty());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_embedding_safetensors_init_from_spec_can_embed() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(QWEN3_EMBEDDING_SAFETENSORS_DIR) {
|
||||
println!("skip safetensors init_from_spec test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: aha::models::WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Safetensors,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some(QWEN3_EMBEDDING_SAFETENSORS_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3EmbeddingModel::init_from_spec(&spec, None, None)?;
|
||||
let output = model.embed(&["test safetensors embedding".to_string()])?;
|
||||
assert_eq!(output.len(), 1);
|
||||
assert!(!output[0].is_empty());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_embedding_load_spec_accepts_gguf() {
|
||||
let spec = LoadSpec {
|
||||
model: aha::models::WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some("D:/model_download/Qwen3-Embedding-0.6B-GGUF/model.gguf".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
spec.validate()
|
||||
.expect("qwen3 embedding should accept gguf artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_embedding_init_from_spec_gguf_can_embed() -> Result<()> {
|
||||
let gguf_path = match first_file_with_extension(QWEN3_EMBEDDING_GGUF_DIR, "gguf") {
|
||||
Ok(path) => path,
|
||||
Err(err) => {
|
||||
println!("skip gguf init_from_spec test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
let spec = LoadSpec {
|
||||
model: aha::models::WhichModel::Qwen3Embedding0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(gguf_path.to_string_lossy().to_string()),
|
||||
tokenizer_dir: Some(QWEN3_EMBEDDING_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3EmbeddingModel::init_from_spec(&spec, None, None)?;
|
||||
let output = model.embed(&["test gguf embedding".to_string()])?;
|
||||
assert_eq!(output.len(), 1);
|
||||
assert!(!output[0].is_empty());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_embedding_real_texts_similarity() -> Result<()> {
|
||||
// Run this test only:
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
use std::{fs, path::Path};
|
||||
|
||||
use aha::models::{
|
||||
ArtifactKind, GenerateModel, LoadSpec, ModelPaths, WhichModel,
|
||||
qwen3::generate::Qwen3GenerateModel,
|
||||
};
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
|
||||
const DEFAULT_QWEN3_SAFETENSORS_DIRS: &[&str] = &[
|
||||
r"D:\model_download\Qwen3-0.6B",
|
||||
r"D:\model_download\Qwen\Qwen3-0.6B",
|
||||
];
|
||||
const DEFAULT_QWEN3_GGUF_FILES: &[&str] = &[
|
||||
r"D:\model_download\Qwen3-0.6B-GGUF\Qwen3-0.6B-Q8_0.gguf",
|
||||
r"D:\model_download\Qwen\Qwen3-0.6B-GGUF\Qwen3-0.6B-Q8_0.gguf",
|
||||
];
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
const DEFAULT_QWEN3_ONNX_DIRS: &[&str] = &[
|
||||
r"D:\model_download\Qwen3-0.6B-ONNX",
|
||||
r"D:\model_download\Qwen\Qwen3-0.6B-ONNX",
|
||||
];
|
||||
|
||||
fn find_first_gguf_file(dir: &Path) -> Option<String> {
|
||||
let entries = fs::read_dir(dir).ok()?;
|
||||
for entry in entries.flatten() {
|
||||
let path = entry.path();
|
||||
if path
|
||||
.extension()
|
||||
.and_then(|ext| ext.to_str())
|
||||
.is_some_and(|ext| ext.eq_ignore_ascii_case("gguf"))
|
||||
{
|
||||
return Some(path.to_string_lossy().to_string());
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn resolve_qwen3_safetensors_dir() -> Option<String> {
|
||||
if let Ok(path) = std::env::var("AHA_QWEN3_SAFETENSORS_DIR")
|
||||
&& Path::new(&path).is_dir()
|
||||
{
|
||||
return Some(path);
|
||||
}
|
||||
|
||||
for path in DEFAULT_QWEN3_SAFETENSORS_DIRS {
|
||||
if Path::new(path).is_dir() {
|
||||
return Some((*path).to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let save_dir = aha::utils::get_default_save_dir()?;
|
||||
let managed_path = format!("{save_dir}/Qwen/Qwen3-0.6B");
|
||||
if Path::new(&managed_path).is_dir() {
|
||||
return Some(managed_path);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn resolve_qwen3_gguf_path() -> Option<String> {
|
||||
if let Ok(path) = std::env::var("AHA_QWEN3_GGUF_PATH")
|
||||
&& Path::new(&path).is_file()
|
||||
{
|
||||
return Some(path);
|
||||
}
|
||||
|
||||
for path in DEFAULT_QWEN3_GGUF_FILES {
|
||||
if Path::new(path).is_file() {
|
||||
return Some((*path).to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let save_dir = aha::utils::get_default_save_dir()?;
|
||||
let managed_dir = format!("{save_dir}/Qwen/Qwen3-0.6B-GGUF");
|
||||
find_first_gguf_file(Path::new(&managed_dir))
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn resolve_qwen3_onnx_dir() -> Option<String> {
|
||||
if let Ok(path) = std::env::var("AHA_QWEN3_ONNX_DIR")
|
||||
&& Path::new(&path).is_dir()
|
||||
{
|
||||
return Some(path);
|
||||
}
|
||||
|
||||
for path in DEFAULT_QWEN3_ONNX_DIRS {
|
||||
if Path::new(path).is_dir() {
|
||||
return Some((*path).to_string());
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn build_text_request() -> Result<ChatCompletionParameters> {
|
||||
let payload = serde_json::json!({
|
||||
"model": "qwen3-0.6b",
|
||||
"max_tokens": 8,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "请用一句话介绍 Rust。"
|
||||
}
|
||||
]
|
||||
});
|
||||
Ok(serde_json::from_value(payload)?)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_safetensors_init_from_spec_can_generate() -> Result<()> {
|
||||
let Some(weight_dir) = resolve_qwen3_safetensors_dir() else {
|
||||
println!("skip qwen3 safetensors test: model dir not found, set AHA_QWEN3_SAFETENSORS_DIR");
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_0_6B,
|
||||
artifact: ArtifactKind::Safetensors,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some(weight_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_load_spec_accepts_gguf() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some("D:/model_download/Qwen3-0.6B-GGUF/model.gguf".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
spec.validate().expect("qwen3 should accept gguf artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_load_spec_accepts_onnx() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some("D:/model_download/Qwen3-0.6B-ONNX".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
spec.validate().expect("qwen3 should accept onnx artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_gguf_init_from_spec_can_generate() -> Result<()> {
|
||||
let Some(gguf_path) = resolve_qwen3_gguf_path() else {
|
||||
println!("skip qwen3 gguf test: model file not found, set AHA_QWEN3_GGUF_PATH");
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(gguf_path),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn qwen3_onnx_init_from_spec_can_generate() -> Result<()> {
|
||||
use aha::models::common::onnx::ensure_ort_dylib_path;
|
||||
|
||||
if let Err(err) = ensure_ort_dylib_path() {
|
||||
println!("skip qwen3 onnx test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let Some(onnx_dir) = resolve_qwen3_onnx_dir() else {
|
||||
println!("skip qwen3 onnx test: onnx dir not found, set AHA_QWEN3_ONNX_DIR");
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3_0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(onnx_dir.clone()),
|
||||
tokenizer_dir: Some(onnx_dir),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3GenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let response = model.generate(build_text_request()?)?;
|
||||
let value = serde_json::to_value(response)?;
|
||||
let choices_len = value
|
||||
.get("choices")
|
||||
.and_then(|choices| choices.as_array())
|
||||
.map_or(0, |choices| choices.len());
|
||||
assert!(choices_len > 0, "expected at least one generated choice");
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
use std::path::Path;
|
||||
|
||||
use aha::models::{
|
||||
ArtifactKind, LoadSpec, ModelPaths, WhichModel, qwen3_reranker::generate::Qwen3RerankerModel,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
|
||||
const QWEN3_RERANKER_ONNX_DIR: &str = r"D:\model_download\Qwen3-Reranker-0.6B-ONNX";
|
||||
const QWEN3_RERANKER_SAFETENSORS_DIR: &str = r"D:\model_download\Qwen3-Reranker-0.6B";
|
||||
const QWEN3_RERANKER_GGUF_DIR: &str = r"D:\model_download\Qwen3-Reranker-0.6B-Q8_0-GGUF";
|
||||
|
||||
fn require_existing_dir(path: &str) -> Result<()> {
|
||||
let dir = Path::new(path);
|
||||
if !dir.exists() {
|
||||
return Err(anyhow!("model dir not found: {}", path));
|
||||
}
|
||||
if !dir.is_dir() {
|
||||
return Err(anyhow!("path is not a directory: {}", path));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn first_file_with_extension(dir: &str, extension: &str) -> Result<String> {
|
||||
require_existing_dir(dir)?;
|
||||
let mut files = std::fs::read_dir(dir)?
|
||||
.flatten()
|
||||
.map(|entry| entry.path())
|
||||
.filter(|path| {
|
||||
path.is_file()
|
||||
&& path
|
||||
.extension()
|
||||
.is_some_and(|ext| ext.eq_ignore_ascii_case(extension))
|
||||
})
|
||||
.map(|path| path.to_string_lossy().to_string())
|
||||
.collect::<Vec<_>>();
|
||||
files.sort();
|
||||
files
|
||||
.into_iter()
|
||||
.next()
|
||||
.ok_or_else(|| anyhow!("no .{} file found in {}", extension, dir))
|
||||
}
|
||||
|
||||
fn top_score_index(scores: &[f32]) -> Option<usize> {
|
||||
scores
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|lhs, rhs| lhs.1.total_cmp(rhs.1))
|
||||
.map(|(idx, _)| idx)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_reranker_onnx_init_from_spec_can_rerank() -> Result<()> {
|
||||
if let Err(err) = aha::models::common::onnx::ensure_ort_dylib_path() {
|
||||
println!("skip reranker onnx test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if let Err(err) = require_existing_dir(QWEN3_RERANKER_ONNX_DIR) {
|
||||
println!("skip reranker onnx test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(QWEN3_RERANKER_ONNX_DIR.to_string()),
|
||||
// Intentionally left None to validate tokenizer auto-fallback from *-ONNX to sibling dir.
|
||||
tokenizer_dir: None,
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3RerankerModel::init_from_spec(&spec, None, None)?;
|
||||
let docs = vec![
|
||||
"Rust async requests are commonly built with reqwest and tokio.".to_string(),
|
||||
"Paris is the capital of France.".to_string(),
|
||||
];
|
||||
let scores = model.rerank("How to make async HTTP calls in Rust?", &docs)?;
|
||||
assert_eq!(scores.len(), docs.len());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_reranker_safetensors_init_from_spec_can_rerank() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(QWEN3_RERANKER_SAFETENSORS_DIR) {
|
||||
println!("skip reranker safetensors test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Safetensors,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some(QWEN3_RERANKER_SAFETENSORS_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut model = Qwen3RerankerModel::init_from_spec(&spec, None, None)?;
|
||||
let docs = vec![
|
||||
"Rust async requests are commonly built with reqwest and tokio.".to_string(),
|
||||
"Paris is the capital of France.".to_string(),
|
||||
];
|
||||
let scores = model.rerank("How to make async HTTP calls in Rust?", &docs)?;
|
||||
assert_eq!(scores.len(), docs.len());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_reranker_load_spec_accepts_gguf() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some("D:/model_download/Qwen3-Reranker-0.6B-GGUF/model.gguf".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
spec.validate()
|
||||
.expect("qwen3 reranker should accept gguf artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_reranker_init_from_spec_gguf_can_rerank() -> Result<()> {
|
||||
let gguf_path = match first_file_with_extension(QWEN3_RERANKER_GGUF_DIR, "gguf") {
|
||||
Ok(path) => path,
|
||||
Err(err) => {
|
||||
println!("skip reranker gguf test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(gguf_path),
|
||||
tokenizer_dir: Some(QWEN3_RERANKER_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let mut model = Qwen3RerankerModel::init_from_spec(&spec, None, None)?;
|
||||
let docs = vec![
|
||||
"Rust async requests are commonly built with reqwest and tokio.".to_string(),
|
||||
"Paris is the capital of France.".to_string(),
|
||||
];
|
||||
let scores = model.rerank("How to make async HTTP calls in Rust?", &docs)?;
|
||||
assert_eq!(scores.len(), docs.len());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn qwen3_reranker_top_doc_consistent_across_formats() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(QWEN3_RERANKER_SAFETENSORS_DIR) {
|
||||
println!("skip reranker consistency test (safetensors): {err}");
|
||||
return Ok(());
|
||||
}
|
||||
if let Err(err) = require_existing_dir(QWEN3_RERANKER_ONNX_DIR) {
|
||||
println!("skip reranker consistency test (onnx): {err}");
|
||||
return Ok(());
|
||||
}
|
||||
if let Err(err) = aha::models::common::onnx::ensure_ort_dylib_path() {
|
||||
println!("skip reranker consistency test (onnxruntime): {err}");
|
||||
return Ok(());
|
||||
}
|
||||
let gguf_path = match first_file_with_extension(QWEN3_RERANKER_GGUF_DIR, "gguf") {
|
||||
Ok(path) => path,
|
||||
Err(err) => {
|
||||
println!("skip reranker consistency test (gguf): {err}");
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
let query = "How to make async HTTP calls in Rust?";
|
||||
let docs = vec![
|
||||
"Rust async requests are commonly built with reqwest and tokio.".to_string(),
|
||||
"Paris is the capital of France.".to_string(),
|
||||
"Database index tuning can improve SQL query speed.".to_string(),
|
||||
];
|
||||
|
||||
let safetensors_spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Safetensors,
|
||||
paths: ModelPaths {
|
||||
weight_dir: Some(QWEN3_RERANKER_SAFETENSORS_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let onnx_spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(QWEN3_RERANKER_ONNX_DIR.to_string()),
|
||||
tokenizer_dir: None,
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let gguf_spec = LoadSpec {
|
||||
model: WhichModel::Qwen3Reranker0_6B,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(gguf_path),
|
||||
tokenizer_dir: Some(QWEN3_RERANKER_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
let mut safetensors_model = Qwen3RerankerModel::init_from_spec(&safetensors_spec, None, None)?;
|
||||
let mut onnx_model = Qwen3RerankerModel::init_from_spec(&onnx_spec, None, None)?;
|
||||
let mut gguf_model = Qwen3RerankerModel::init_from_spec(&gguf_spec, None, None)?;
|
||||
|
||||
let safetensors_scores = safetensors_model.rerank(query, &docs)?;
|
||||
let onnx_scores = onnx_model.rerank(query, &docs)?;
|
||||
let gguf_scores = gguf_model.rerank(query, &docs)?;
|
||||
|
||||
let safetensors_top = top_score_index(&safetensors_scores)
|
||||
.ok_or_else(|| anyhow!("safetensors scores are empty"))?;
|
||||
let onnx_top = top_score_index(&onnx_scores).ok_or_else(|| anyhow!("onnx scores are empty"))?;
|
||||
let gguf_top = top_score_index(&gguf_scores).ok_or_else(|| anyhow!("gguf scores are empty"))?;
|
||||
|
||||
assert_eq!(
|
||||
safetensors_top, 0,
|
||||
"expected safetensors top doc to be rust doc"
|
||||
);
|
||||
assert_eq!(onnx_top, 0, "expected onnx top doc to be rust doc");
|
||||
assert_eq!(gguf_top, 0, "expected gguf top doc to be rust doc");
|
||||
Ok(())
|
||||
}
|
||||
Reference in New Issue
Block a user