添加 rerank 和 embedding模型支持,并且添加 onnx 模型

This commit is contained in:
273265088@qq.com
2026-03-24 16:18:39 +08:00
parent 77b244e53e
commit b8d8732c3d
28 changed files with 2130 additions and 158 deletions
Generated
+60
View File
@@ -42,6 +42,7 @@ dependencies = [
"minijinja",
"modelscope",
"num",
"ort",
"rayon",
"realfft",
"reqwest 0.12.28",
@@ -2710,6 +2711,16 @@ dependencies = [
"regex-automata",
]
[[package]]
name = "matrixmultiply"
version = "0.3.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a06de3016e9fae57a36fd14dba131fccf49f74b40b7fbdb472f96e361ec71a08"
dependencies = [
"autocfg",
"rawpointer",
]
[[package]]
name = "maybe-rayon"
version = "0.1.1"
@@ -2887,6 +2898,21 @@ dependencies = [
"tempfile",
]
[[package]]
name = "ndarray"
version = "0.17.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "520080814a7a6b4a6e9070823bb24b4531daac8c4627e08ba5de8c5ef2f2752d"
dependencies = [
"matrixmultiply",
"num-complex",
"num-integer",
"num-traits",
"portable-atomic",
"portable-atomic-util",
"rawpointer",
]
[[package]]
name = "new_debug_unreachable"
version = "1.0.6"
@@ -3184,6 +3210,25 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "04744f49eae99ab78e0d5c0b603ab218f515ea8cfe5a456d7629ad883a3b6e7d"
[[package]]
name = "ort"
version = "2.0.0-rc.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7de3af33d24a745ffb8fab904b13478438d1cd52868e6f17735ef6e1f8bf133"
dependencies = [
"libloading 0.9.0",
"ndarray",
"ort-sys",
"smallvec",
"tracing",
]
[[package]]
name = "ort-sys"
version = "2.0.0-rc.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7b497d21a8b6fbb4b5a544f8fadb77e801a09ae0add9e411d31c6f89e3c1e90"
[[package]]
name = "parking_lot"
version = "0.12.5"
@@ -3295,6 +3340,15 @@ version = "1.13.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49"
[[package]]
name = "portable-atomic-util"
version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "091397be61a01d4be58e7841595bd4bfedb15f1cd54977d79b8271e94ed799a3"
dependencies = [
"portable-atomic",
]
[[package]]
name = "potential_utf"
version = "0.1.4"
@@ -3687,6 +3741,12 @@ dependencies = [
"bitflags 2.11.0",
]
[[package]]
name = "rawpointer"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "60a357793950651c4ed0f3f52338f53b2f809f32d83a07f72909fa13e4c6c1e3"
[[package]]
name = "rayon"
version = "1.11.0"
+3
View File
@@ -54,3 +54,6 @@ ffmpeg = ["ffmpeg-next"]
needless_range_loop = "allow"
single_range_in_vec_init = "allow"
manual_div_ceil = "allow"
[dev-dependencies]
ort = { version = "2.0.0-rc.10", default-features = false, features = ["load-dynamic", "api-24"] }
+547
View File
@@ -0,0 +1,547 @@
# AHA 开发规范
本文档不是通用 Rust 模板,而是基于当前 `aha` 仓库真实代码结构整理出的项目级规范。目标有两个:
1. 统一现有模型目录的职责分层。
2. 为未来新增原生模型、GGUF 模型、ONNX 模型提供可执行的接入标准。
---
## 1. 当前项目真实结构
### 1.1 顶层目录职责
```text
aha/
├── src/
│ ├── main.rs # CLI 与服务启动入口
│ ├── lib.rs # 对外库导出
│ ├── api/ # HTTP/OpenAI 兼容接口
│ ├── exec/ # `aha run` 的模型直调入口
│ ├── models/ # 模型实现与模型工厂
│ ├── tokenizer/ # tokenizer 封装
│ ├── chat_template/ # chat template 封装
│ ├── position_embed/ # 位置编码
│ ├── process.rs # 服务进程管理
│ └── utils/ # 下载、设备、dtype、文件查找等通用工具
├── tests/ # 集成测试与格式测试
├── docs/ # 面向用户的文档
└── dev_rule.md # 本规范
```
### 1.2 模型接入链路
新增一个模型,必须理解下面这条链路,而不是只加 `src/models/<name>/`
1. `src/models/<name>/`
实现模型本体、配置、推理入口。
2. `src/models/mod.rs`
注册模块、枚举 `WhichModel`、模型元信息、模型工厂 `load_model``ModelInstance` 能力分派。
3. `src/exec/<name>.rs`
接入 `aha run -m <model>` 的直调能力。
4. `src/exec/mod.rs`
导出 exec 模块。
5. `src/main.rs`
`WhichModel` 路由到具体 exec 实现;必要时处理 GGUF / mmproj / ONNX 的路径参数。
6. `src/api/mod.rs`
若模型属于 `chat` / `embedding` / `rerank` / `asr` / `ocr` 等服务能力,需要保证 `ModelInstance` 对应方法可用。
7. `tests/`
至少补齐加载测试和最小推理测试。
8. `docs/``README*``changelog`
更新用户可见说明。
---
## 2. 模型目录统一规范
### 2.1 最低标准
从现在开始,`src/models/<model_name>/` 的最低结构如下:
```text
src/models/<model_name>/
├── mod.rs
├── config.rs
├── model.rs
└── generate.rs
```
说明:
- `mod.rs` 必须只做模块导出,不写业务逻辑。
- `config.rs` 负责配置结构体、默认值、配置文件加载辅助。
- `model.rs` 负责底层网络结构、后端封装、权重加载、核心推理逻辑。
- `generate.rs` 负责对外可调用入口。
### 2.2 可选文件
满足以下条件时再新增文件,不要随意拆目录:
- `processor.rs`
仅用于多模态输入预处理或复杂后处理,例如图像、音频、视频、OCR patch 处理。
- `tokenizer.rs`
仅当模型目录需要自定义 tokenizer 适配,且通用 `TokenizerModel` 不够用。
- `audio_vae.rs``minicpm4.rs` 等附属文件
仅当模型实现包含稳定、独立的子模块时允许拆出。
### 2.3 各文件职责边界
#### `mod.rs`
只允许:
- `pub mod ...`
- 必要时 `pub use ...`
不允许:
- 业务实现
- 大段常量
- 初始化逻辑
#### `config.rs`
负责:
- `config.json`
- `generation_config.json`
- `preprocessor_config.json`
- 与格式相关但不属于运行态的配置映射
不负责:
- 真正加载权重
- 调用 tokenizer
- 调用 Candle forward
#### `model.rs`
负责:
- Candle 模型结构体
- 权重加载
- 格式后端封装
- 核心 `forward` / `embed` / `rerank` / `transcribe` / `process_image` 等内部能力
不负责:
- CLI 参数解析
- HTTP 请求结构
- `aha_openai_dive` 响应组装
#### `generate.rs`
负责:
- 作为当前模型目录的对外入口类型
- 将 tokenizer、processor、model backend 串起来
- 实现 `GenerateModel` trait,或提供该模型暴露给 `ModelInstance` 的统一入口
约束:
- 对外入口类型名必须稳定,供 `src/models/mod.rs``src/exec/*``tests/*` 使用。
- 目录重构时优先保持公开方法不变,例如 `init``generate``generate_stream``embed``rerank`
---
## 3. 按能力划分的目录模板
### 3.1 文本生成模型
适用:
- Qwen3
- Qwen3.5
- MiniCPM4
模板:
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
└── generate.rs
```
规则:
- `generate.rs` 负责 `GenerateModel` 实现。
- `model.rs` 负责 `forward`、KV cache、权重加载。
### 3.2 多模态模型
适用:
- Qwen3VL
- Qwen2.5VL
- Qwen3ASR
- DeepSeek-OCR
- Hunyuan-OCR
- PaddleOCR-VL
模板:
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
├── generate.rs
└── processor.rs
```
规则:
- `processor.rs` 负责把请求输入变成模型可消费的 tensor 或替换文本。
- `generate.rs` 只做流程编排,不堆复杂前处理。
### 3.3 Embedding / Reranker 模型
适用:
- Qwen3-Embedding
- Qwen3-Reranker
- 后续任意检索相关模型
模板:
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
└── generate.rs
```
规则:
- `model.rs` 负责 embedding backend 或 rerank backend。
- `generate.rs` 提供稳定对外入口,供 `ModelInstance::embedding` / `ModelInstance::rerank` 调用。
- 公共检索算法优先放在 `src/models/common/retrieval.rs`,不要在各模型目录重复实现。
### 3.4 纯图像或工具型模型
适用:
- RMBG2.0
模板:
```text
src/models/<model>/
├── mod.rs
├── model.rs
└── generate.rs
```
补充要求:
- 如果模型已经存在稳定配置文件,仍然建议补 `config.rs`
- 新增此类模型时,不要因为现有个别旧目录较薄,就继续复制旧写法。
---
## 4. 多格式模型规范
### 4.1 格式定义
当前项目已经在 `src/models/mod.rs` 中抽象出:
- `Safetensors`
- `Gguf`
- `Onnx`
新增模型时,必须先明确自己属于哪一种格式,不能等实现完成后再补枚举。
### 4.2 格式接入原则
#### 原生 / Safetensors 模型
规则:
- 默认走项目当前主路径。
- 权重查找统一使用 `find_type_files(path, "safetensors")`
- 设备与 dtype 统一通过 `get_device``get_dtype` 获取。
- 能被自动下载的模型,必须在 `WhichModel::is_download_managed()` 语义下成立。
#### GGUF 模型
规则:
- 必须在 `WhichModel::artifact_format()` 返回 `ModelArtifactFormat::Gguf`
- `main.rs` 中必须要求 `--gguf-path`,多模态 GGUF 额外支持 `--mmproj-path`
- tokenizer 若从 GGUF metadata 构建,必须封装在 `model.rs` 或专门 backend 中,不允许散落在 exec 或 main。
- GGUF 的 mmproj 仅属于后端实现细节,不应污染上层 API 类型。
#### ONNX 模型
规则:
- 必须在 `WhichModel::artifact_format()` 返回 `ModelArtifactFormat::Onnx`
- 若运行时尚未集成 ONNX 推理,仍然要保证:
- 模型枚举和格式分类准确;
- 测试能够验证 ONNX 文件可发现、可读取、可创建 session;
- 错误提示明确说明“格式已识别,但 runtime 未集成”。
- ONNX Runtime 的 session 创建、输入输出张量映射必须收敛到模型目录内部,不允许堆到测试文件里作为“事实实现”。
### 4.3 多格式目录组织建议
对未来一个同时支持 `safetensors + gguf + onnx` 的模型,推荐结构:
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
├── generate.rs
├── backend_safetensors.rs # 当实现明显变大时再拆
├── backend_gguf.rs
└── backend_onnx.rs
```
约束:
- 如果格式实现还很轻,先放在 `model.rs` 的私有 backend 结构体里。
- 只有当 `model.rs` 已明显过大,才允许拆 `backend_*` 文件。
- 不允许一开始就为了“看起来完整”拆很多空文件。
### 4.4 多格式统一入口要求
对于同一模型家族,不管底层格式是什么,对外入口必须统一:
- 同一个 `WhichModel` 语义只表达“模型变体”,不表达内部临时实现。
- `generate.rs` 的公开类型应稳定。
- `src/models/mod.rs::load_model()` 负责根据 `WhichModel` 选择正确初始化路径。
---
## 5. 命名与导出规范
### 5.1 目录命名
- 目录名统一使用小写蛇形命名。
- 与上游模型 ID 不一致时,以项目内部统一命名为准,例如 `qwen3_asr``qwen3_embedding`
### 5.2 类型命名
- 配置类型:`<ModelName>Config`
- 生成配置类型:`<ModelName>GenerationConfig`
- 底层模型:`<ModelName>Model``<ModelName>Backend`
- 对外入口:`<ModelName>GenerateModel``<ModelName>EmbeddingModel``<ModelName>RerankerModel`
- 处理器:`<ModelName>Processor`
### 5.3 `mod.rs` 命名要求
必须导出到以下位置:
- `src/models/mod.rs`
- `src/exec/mod.rs`
禁止:
- 同一个模型目录里同时存在多个对外主入口但没有清晰职责说明。
---
## 6. 新增模型的强制清单
### 6.1 模型目录
必须完成:
1. 创建 `src/models/<model>/`
2. 按模板补齐最少文件
3. 保持职责边界清晰
### 6.2 模型注册
必须同时修改 `src/models/mod.rs`
1. `pub mod <model>;`
2. `WhichModel` 新增枚举项
3. `LISTED_MODELS` 增加公开模型
4. `artifact_format()`
5. `openai_model_id()`
6. `owner()`
7. `model_id()`
8. `model_type()`
9. `ModelInstance` 增加实例变体
10. `load_model()` 增加加载分支
### 6.3 CLI 接入
必须同时修改:
1. `src/exec/<model>.rs`
2. `src/exec/mod.rs`
3. `src/main.rs`
要求:
- `aha run` 必须能走到该模型。
- GGUF / ONNX 如有特殊路径参数,必须在 `main.rs` 路由时处理。
### 6.4 API 接入
当模型支持服务能力时,必须确保:
- `chat` 类型实现 `GenerateModel`
- `embedding` 类型实现 `ModelInstance::embedding`
- `rerank` 类型实现 `ModelInstance::rerank`
- `asr` / `ocr` / `image` 类型能复用现有 API 行为
### 6.5 文档接入
至少更新:
- `README.md`
- `README.zh-CN.md`
- `docs/supported-models.md`
- `docs/supported-models.zh-CN.md`
- `docs/changelog.md`
- `docs/changelog.zh-CN.md`
如果 CLI 或 API 行为变更,还要更新:
- `docs/cli*.md`
- `docs/api*.md`
---
## 7. 测试规范
### 7.1 每个新模型至少要有的测试
1. 加载测试
验证模型文件能被发现并初始化。
2. 最小推理测试
验证 `generate` / `embed` / `rerank` / `transcribe` / `remove_background` 至少能跑通一次。
3. 错误输入测试
验证空输入、非法输入、缺失文件路径能返回明确错误。
### 7.2 多格式模型测试矩阵
对于 `safetensors + gguf + onnx` 多格式模型,建议至少覆盖:
1. safetensors 加载测试
2. gguf 文件发现与 tokenizer/session/backend 加载测试
3. onnx 文件发现测试
4. onnxruntime session 创建测试
5. 至少一种真实输入的结果合理性测试
### 7.3 测试文件命名
- 单模型:`tests/test_<model>.rs`
- 多格式专项:`tests/test_<model>_multi_format.rs`
### 7.4 校验命令
最少校验:
```bash
cargo fmt --check
cargo check
```
建议校验:
```bash
cargo test --test <target_test>
```
注意:
- 大模型测试通常依赖本地权重,不要求在所有环境全量跑通。
- 但编译检查必须通过。
---
## 8. 代码风格约束
### 8.1 不允许的做法
-`mod.rs` 堆业务逻辑
-`exec` 中写模型核心推理代码
-`main.rs` 中写模型权重加载细节
- 在测试文件里藏正式实现逻辑
- 为了“结构完整”创建一堆空模块
- 复制已有模型代码后只改字符串,不抽象公共逻辑
### 8.2 推荐做法
- 公共检索逻辑放 `src/models/common/retrieval.rs`
- 公共 GGUF 工具放 `src/models/common/gguf.rs`
- 设备、dtype、权重文件发现统一走 `src/utils/mod.rs`
- 公开方法保持小而稳定,内部再分层
### 8.3 兼容性原则
当重构模型目录时:
- 优先保持对外类型名不变
- 优先保持 `init(...)` 签名不变
- 优先保持 `src/models/mod.rs``src/exec/*`、测试里的调用方式不变
如果必须改公开接口,必须同步修改:
- `src/models/mod.rs`
- `src/exec/*`
- `src/main.rs`
- `src/api/mod.rs`
- `tests/*`
- 相关文档
---
## 9. 当前仓库的落地结论
根据当前仓库实际情况,后续新增模型必须遵循以下结论:
1. `qwen3_embedding``qwen3_reranker` 现在也必须按 `config.rs + model.rs + generate.rs + mod.rs` 维护。
2. 未来新增 Embedding/Reranker 模型,不允许再只保留 `mod.rs + generate.rs`
3. `docs/development.md` 当前更接近通用模板,不应作为新增模型时的唯一依据;以本文件和实际代码结构为准。
4. ONNX 目前在仓库里已经有测试与格式识别语义,但运行时集成尚未完整落地;新增 ONNX 模型时,必须同时补 runtime 设计,而不是只补测试。
5. GGUF 模型接入必须从一开始就考虑 `main.rs` 路径参数、`WhichModel::artifact_format()``load_model()` 的完整链路。
---
## 10. 新增模型模板清单
### 10.1 原生文本模型
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
└── generate.rs
```
### 10.2 原生多模态模型
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
├── generate.rs
└── processor.rs
```
### 10.3 多格式检索模型
```text
src/models/<model>/
├── mod.rs
├── config.rs
├── model.rs
├── generate.rs
├── backend_safetensors.rs # 可选
├── backend_gguf.rs # 可选
└── backend_onnx.rs # 可选
```
最终原则只有一句:
新增模型不是“把代码放进一个目录”,而是“把模型完整接入到 models factory、CLI、API、tests、docs 五条链路中”,并且目录职责必须与现有标准结构保持一致。
+80
View File
@@ -9,6 +9,22 @@ aha supports a growing collection of state-of-the-art AI models across multiple
| **Qwen3-0.6B** | 0.6B | Latest generation | Advanced reasoning | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **MiniCPM4-0.5B** | 0.5B | Efficient lightweight | Edge deployment | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Embedding
| Model | Parameters | Description | License |
|-------|-----------|-------------|---------|
| **Qwen3-Embedding-0.6B** | 0.6B | Text embedding (safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Embedding-4B** | 4B | Text embedding (safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Embedding-8B** | 8B | Text embedding (safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Reranker
| Model | Parameters | Description | License |
|-------|-----------|-------------|---------|
| **Qwen3-Reranker-0.6B** | 0.6B | Text reranking (embedding-similarity baseline, safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Reranker-4B** | 4B | Text reranking (embedding-similarity baseline, safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Reranker-8B** | 8B | Text reranking (embedding-similarity baseline, safetensors) | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Vision & Multimodal
| Model | Parameters | Description | License |
@@ -23,6 +39,18 @@ aha supports a growing collection of state-of-the-art AI models across multiple
| **Qwen3.5-2B** | 2B | Native Multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-4B** | 4B | Native Multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-9B** | 9B | Native Multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2** | 9B | Distilled variant (Qwen3.5 family) | [Model license on HF](https://huggingface.co/Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2) |
### Qwen3.5 GGUF Sources (Runtime Reused)
- Jackrong/Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF
- Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF
- unsloth/Qwen3.5-0.8B-GGUF
- unsloth/Qwen3.5-2B-GGUF
- unsloth/Qwen3.5-4B-GGUF
- lmstudio-community/Qwen3.5-0.8B-GGUF
- lmstudio-community/Qwen3.5-2B-GGUF
- lmstudio-community/Qwen3.5-4B-GGUF
## OCR
@@ -64,6 +92,58 @@ Models are sourced from:
- [Hugging Face](https://huggingface.co) - Primary model hub
- [ModelScope](https://modelscope.cn) - Chinese model hub
## Registered Repositories (Not Runtime-Integrated Yet)
The following repositories are now cataloged for future integration, but are **not** directly runnable in current `aha` runtime yet:
### MLX / Format-Specific Variants
- Jackrong/MLX-Qwen3.5-27B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-6bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-8bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-6bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-8bit
### Embedding Models
- google/embeddinggemma-300m
- ggml-org/embeddinggemma-300M-GGUF
- onnx-community/embeddinggemma-300m-ONNX
- unsloth/embeddinggemma-300m-GGUF
- onnx-community/Qwen3-Embedding-0.6B-ONNX
- Qwen/Qwen3-Embedding-0.6B-GGUF
- onnx-community/Qwen3-Embedding-4B-ONNX
- Qwen/Qwen3-Embedding-4B-GGUF
- Qwen/Qwen3-Embedding-8B-GGUF
- onnx-community/Qwen3-Embedding-8B-ONNX
- perplexity-ai/pplx-embed-v1-0.6b
- nomic-ai/nomic-embed-text-v2-moe
- nomic-ai/nomic-embed-text-v2-moe-GGUF
- jinaai/jina-embeddings-v5-text-small
- jinaai/jina-embeddings-v5-text-nano
- jinaai/jina-embeddings-v5-text-small-text-matching
- jinaai/jina-embeddings-v5-text-small-text-matching-GGUF
- sentence-transformers/all-MiniLM-L6-v2
- sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2
### Reranker Models
- BAAI/bge-reranker-v2-m3
- ggml-org/Qwen3-Reranker-0.6B-Q8_0-GGUF
### ONNX Repositories
- onnx-community/GLM-OCR-ONNX
- onnx-community/Qwen3-Reranker-0.6B-ONNX
- onnx-community/Qwen3.5-2B-ONNX
- onnx-community/Qwen3.5-4B-ONNX
- onnx-community/Qwen3.5-0.8B-ONNX
- onnx-community/Qwen3-VL-2B-Instruct-ONNX
- onnx-community/ONNX_Qwen3-Embedding-0.6B
- onnx-community/Nanbeige4.1-3B-ONNX
- onnx-community/Qwen3-Embedding-8B-ONNX
- onnx-community/Qwen3-Embedding-4B-ONNX
- onnx-community/bge-reranker-v2-m3-ONNX
- onnx-community/all-MiniLM-L6-v2-ONNX
## Adding New Models
See [Development Guide](./development.md) for instructions on adding new model integrations.
+80
View File
@@ -9,6 +9,22 @@ aha 支持多个领域的最先进 AI 模型集合。
| **Qwen3-0.6B** | 0.6B | 最新一代 | 高级推理 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **MiniCPM4-0.5B** | 0.5B | 高效轻量级 | 边缘部署 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Embedding
| 模型 | 参数量 | 描述 | 开源协议 |
|------|--------|------|---------|
| **Qwen3-Embedding-0.6B** | 0.6B | 文本向量(safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Embedding-4B** | 4B | 文本向量(safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Embedding-8B** | 8B | 文本向量(safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## Reranker
| 模型 | 参数量 | 描述 | 开源协议 |
|------|--------|------|---------|
| **Qwen3-Reranker-0.6B** | 0.6B | 文本重排(基于 embedding 相似度的基线实现,safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Reranker-4B** | 4B | 文本重排(基于 embedding 相似度的基线实现,safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3-Reranker-8B** | 8B | 文本重排(基于 embedding 相似度的基线实现,safetensors | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
## 视觉与多模态
| 模型 | 参数量 | 描述 | 开源协议 |
@@ -23,6 +39,18 @@ aha 支持多个领域的最先进 AI 模型集合。
| **Qwen3.5-2B** | 2B | 原生多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-4B** | 4B | 原生多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-9B** | 9B | 原生多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2** | 9B | Qwen3.5 同构蒸馏版本 | [HF 页面许可证](https://huggingface.co/Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2) |
### Qwen3.5 GGUF 仓库来源(复用现有运行时)
- Jackrong/Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF
- Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF
- unsloth/Qwen3.5-0.8B-GGUF
- unsloth/Qwen3.5-2B-GGUF
- unsloth/Qwen3.5-4B-GGUF
- lmstudio-community/Qwen3.5-0.8B-GGUF
- lmstudio-community/Qwen3.5-2B-GGUF
- lmstudio-community/Qwen3.5-4B-GGUF
## OCR
@@ -64,6 +92,58 @@ aha 支持多个领域的最先进 AI 模型集合。
- [Hugging Face](https://huggingface.co) - 主模型中心
- [ModelScope](https://modelscope.cn) - 中文模型中心
## 已收录仓库(当前运行时暂未直接接入)
以下仓库已纳入项目模型目录,但当前 `aha` 运行时尚不能直接推理:
### MLX / 特定格式变体
- Jackrong/MLX-Qwen3.5-27B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-6bit
- Jackrong/MLX-Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-8bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-4bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-6bit
- Jackrong/MLX-Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-8bit
### Embedding 模型
- google/embeddinggemma-300m
- ggml-org/embeddinggemma-300M-GGUF
- onnx-community/embeddinggemma-300m-ONNX
- unsloth/embeddinggemma-300m-GGUF
- onnx-community/Qwen3-Embedding-0.6B-ONNX
- Qwen/Qwen3-Embedding-0.6B-GGUF
- onnx-community/Qwen3-Embedding-4B-ONNX
- Qwen/Qwen3-Embedding-4B-GGUF
- Qwen/Qwen3-Embedding-8B-GGUF
- onnx-community/Qwen3-Embedding-8B-ONNX
- perplexity-ai/pplx-embed-v1-0.6b
- nomic-ai/nomic-embed-text-v2-moe
- nomic-ai/nomic-embed-text-v2-moe-GGUF
- jinaai/jina-embeddings-v5-text-small
- jinaai/jina-embeddings-v5-text-nano
- jinaai/jina-embeddings-v5-text-small-text-matching
- jinaai/jina-embeddings-v5-text-small-text-matching-GGUF
- sentence-transformers/all-MiniLM-L6-v2
- sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2
### Reranker 模型
- BAAI/bge-reranker-v2-m3
- ggml-org/Qwen3-Reranker-0.6B-Q8_0-GGUF
### ONNX 仓库
- onnx-community/GLM-OCR-ONNX
- onnx-community/Qwen3-Reranker-0.6B-ONNX
- onnx-community/Qwen3.5-2B-ONNX
- onnx-community/Qwen3.5-4B-ONNX
- onnx-community/Qwen3.5-0.8B-ONNX
- onnx-community/Qwen3-VL-2B-Instruct-ONNX
- onnx-community/ONNX_Qwen3-Embedding-0.6B
- onnx-community/Nanbeige4.1-3B-ONNX
- onnx-community/Qwen3-Embedding-8B-ONNX
- onnx-community/Qwen3-Embedding-4B-ONNX
- onnx-community/bge-reranker-v2-m3-ONNX
- onnx-community/all-MiniLM-L6-v2-ONNX
## 添加新模型
参见 [开发指南](./development.zh-CN.md) 了解添加新模型集成的说明。
+12
View File
@@ -18,6 +18,9 @@ show_help() {
echo " qwen2.5vl-3b"
echo " qwen2.5vl-7b"
echo " qwen3-0.6b"
echo " qwen3-reranker-0.6b"
echo " qwen3-reranker-4b"
echo " qwen3-reranker-8b"
echo " qwen3asr-0.6b"
echo " qwen3asr-1.7b"
echo " qwen3vl-2b"
@@ -58,6 +61,15 @@ case $MODEL_ALIAS in
"qwen3-0.6b")
MODEL_ID="Qwen/Qwen3-0.6B"
;;
"qwen3-reranker-0.6b")
MODEL_ID="Qwen/Qwen3-Reranker-0.6B"
;;
"qwen3-reranker-4b")
MODEL_ID="Qwen/Qwen3-Reranker-4B"
;;
"qwen3-reranker-8b")
MODEL_ID="Qwen/Qwen3-Reranker-8B"
;;
"qwen3asr-0.6b")
MODEL_ID="Qwen/Qwen3-ASR-0.6B"
;;
+267 -80
View File
@@ -7,7 +7,7 @@ use aha::process::cleanup_pid_file;
use aha::utils::string_to_static_str;
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use rocket::futures::StreamExt;
use rocket::serde::{Serialize, json::Json};
use rocket::serde::{Deserialize, Serialize, json::Json};
use rocket::{
Request, State,
futures::Stream,
@@ -16,6 +16,7 @@ use rocket::{
post,
response::{Responder, stream::TextStream},
};
use serde_json::Value;
use tokio::sync::RwLock;
// ASR (Automatic Speech Recognition) API module
@@ -191,6 +192,191 @@ pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, Stri
}
}
#[derive(Debug, Deserialize)]
pub(crate) struct EmbeddingRequest {
pub model: Option<String>,
pub input: Value,
}
#[derive(Debug, Serialize)]
struct EmbeddingData {
object: String,
index: usize,
embedding: Vec<f32>,
}
#[derive(Debug, Serialize)]
struct EmbeddingResponse {
object: String,
data: Vec<EmbeddingData>,
model: String,
}
#[derive(Debug, Deserialize)]
pub(crate) struct RerankRequest {
pub model: Option<String>,
pub query: String,
pub documents: Vec<String>,
pub top_n: Option<usize>,
}
#[derive(Debug, Serialize)]
struct RerankResult {
index: usize,
relevance_score: f32,
document: String,
}
#[derive(Debug, Serialize)]
struct RerankResponse {
object: String,
model: String,
results: Vec<RerankResult>,
}
fn parse_embedding_input(input: &Value) -> anyhow::Result<Vec<String>> {
match input {
Value::String(s) => Ok(vec![s.clone()]),
Value::Array(arr) => {
let mut out = Vec::with_capacity(arr.len());
for v in arr {
let s = v.as_str().ok_or_else(|| {
anyhow::anyhow!("embedding input array must contain only strings")
})?;
out.push(s.to_string());
}
if out.is_empty() {
return Err(anyhow::anyhow!("embedding input cannot be empty"));
}
Ok(out)
}
_ => Err(anyhow::anyhow!(
"embedding input must be a string or an array of strings"
)),
}
}
fn validate_rerank_input(query: &str, documents: &[String]) -> anyhow::Result<()> {
if query.trim().is_empty() {
return Err(anyhow::anyhow!("rerank query cannot be empty"));
}
if documents.is_empty() {
return Err(anyhow::anyhow!("rerank documents cannot be empty"));
}
if documents.iter().any(|doc| doc.trim().is_empty()) {
return Err(anyhow::anyhow!(
"rerank documents cannot contain empty strings"
));
}
Ok(())
}
#[post("/embeddings", data = "<req>")]
pub(crate) async fn embeddings(req: Json<EmbeddingRequest>) -> (Status, Json<Value>) {
let texts = match parse_embedding_input(&req.input) {
Ok(v) => v,
Err(e) => {
return (
Status::BadRequest,
Json(serde_json::json!({ "error": e.to_string() })),
);
}
};
let model_ref = match MODEL.get().cloned() {
Some(v) => v,
None => {
return (
Status::ServiceUnavailable,
Json(serde_json::json!({ "error": "model not init" })),
);
}
};
let mut guard = model_ref.write().await;
let embeddings = match guard.instance.embedding(&texts) {
Ok(v) => v,
Err(e) => {
return (
Status::BadRequest,
Json(serde_json::json!({ "error": e.to_string() })),
);
}
};
let model_name = req
.model
.clone()
.unwrap_or_else(|| guard.which_model.openai_model_id().to_string());
let data = embeddings
.into_iter()
.enumerate()
.map(|(index, embedding)| EmbeddingData {
object: "embedding".to_string(),
index,
embedding,
})
.collect::<Vec<_>>();
let response = EmbeddingResponse {
object: "list".to_string(),
data,
model: model_name,
};
(Status::Ok, Json(serde_json::to_value(response).unwrap()))
}
#[post("/rerank", data = "<req>")]
pub(crate) async fn rerank(req: Json<RerankRequest>) -> (Status, Json<Value>) {
let req = req.into_inner();
if let Err(e) = validate_rerank_input(&req.query, &req.documents) {
return (
Status::BadRequest,
Json(serde_json::json!({ "error": e.to_string() })),
);
}
let model_ref = match MODEL.get().cloned() {
Some(v) => v,
None => {
return (
Status::ServiceUnavailable,
Json(serde_json::json!({ "error": "model not init" })),
);
}
};
let mut guard = model_ref.write().await;
let scores = match guard.instance.rerank(&req.query, &req.documents) {
Ok(v) => v,
Err(e) => {
return (
Status::BadRequest,
Json(serde_json::json!({ "error": e.to_string() })),
);
}
};
let mut results = scores
.into_iter()
.enumerate()
.map(|(index, relevance_score)| RerankResult {
index,
relevance_score,
document: req.documents[index].clone(),
})
.collect::<Vec<_>>();
results.sort_by(|a, b| b.relevance_score.total_cmp(&a.relevance_score));
if let Some(top_n) = req.top_n {
results.truncate(top_n.min(results.len()));
}
let response = RerankResponse {
object: "list".to_string(),
model: req
.model
.unwrap_or_else(|| guard.which_model.openai_model_id().to_string()),
results,
};
(Status::Ok, Json(serde_json::to_value(response).unwrap()))
}
// Health check endpoint
#[derive(Serialize)]
@@ -243,63 +429,6 @@ struct ErrorResponse {
error: String,
}
/// Convert WhichModel to a display-friendly model ID (kebab-case)
fn which_model_to_id(which_model: WhichModel) -> &'static str {
match which_model {
WhichModel::MiniCPM4_0_5B => "minicpm4-0.5b",
WhichModel::Qwen2_5vl3B => "qwen2.5vl-3b",
WhichModel::Qwen2_5vl7B => "qwen2.5vl-7b",
WhichModel::Qwen3_0_6B => "qwen3-0.6b",
WhichModel::Qwen3_5_0_8B => "qwen3.5-0.8b",
WhichModel::Qwen3_5_2B => "qwen3.5-2b",
WhichModel::Qwen3_5_4B => "qwen3.5-4b",
WhichModel::Qwen3_5_9B => "qwen3.5-9b",
WhichModel::Qwen3_5Gguf => "qwen3.5-gguf",
WhichModel::Qwen3ASR0_6B => "qwen3asr-0.6b",
WhichModel::Qwen3ASR1_7B => "qwen3asr-1.7b",
WhichModel::Qwen3vl2B => "qwen3vl-2b",
WhichModel::Qwen3vl4B => "qwen3vl-4b",
WhichModel::Qwen3vl8B => "qwen3vl-8b",
WhichModel::Qwen3vl32B => "qwen3vl-32b",
WhichModel::DeepSeekOCR => "deepseek-ocr",
WhichModel::DeepSeekOCR2 => "deepseek-ocr2",
WhichModel::HunyuanOCR => "hunyuan-ocr",
WhichModel::PaddleOCRVL => "paddleocr-vl",
WhichModel::PaddleOCRVL1_5 => "paddleocr-vl1.5",
WhichModel::RMBG2_0 => "rmbg2.0",
WhichModel::VoxCPM => "voxcpm",
WhichModel::VoxCPM1_5 => "voxcpm1.5",
WhichModel::GlmASRNano2512 => "glm-asr-nano-2512",
WhichModel::FunASRNano2512 => "fun-asr-nano-2512",
WhichModel::GlmOCR => "glm-ocr",
}
}
/// Get the owner/organization name for a model
fn which_model_to_owner(which_model: WhichModel) -> &'static str {
match which_model {
WhichModel::MiniCPM4_0_5B => "OpenBMB",
WhichModel::Qwen2_5vl3B | WhichModel::Qwen2_5vl7B => "Qwen",
WhichModel::Qwen3_0_6B | WhichModel::Qwen3ASR0_6B | WhichModel::Qwen3ASR1_7B => "Qwen",
WhichModel::Qwen3vl2B
| WhichModel::Qwen3vl4B
| WhichModel::Qwen3vl8B
| WhichModel::Qwen3vl32B
| WhichModel::Qwen3_5Gguf => "Qwen",
WhichModel::Qwen3_5_0_8B
| WhichModel::Qwen3_5_2B
| WhichModel::Qwen3_5_4B
| WhichModel::Qwen3_5_9B => "Qwen",
WhichModel::DeepSeekOCR | WhichModel::DeepSeekOCR2 => "deepseek-ai",
WhichModel::HunyuanOCR => "Tencent-Hunyuan",
WhichModel::PaddleOCRVL | WhichModel::PaddleOCRVL1_5 => "PaddlePaddle",
WhichModel::RMBG2_0 => "AI-ModelScope",
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "OpenBMB",
WhichModel::GlmASRNano2512 | WhichModel::GlmOCR => "ZhipuAI",
WhichModel::FunASRNano2512 => "FunAudioLLM",
}
}
#[get("/models")]
pub(crate) async fn models() -> (Status, (ContentType, Json<serde_json::Value>)) {
if let Some(model_ref) = MODEL.get() {
@@ -307,10 +436,10 @@ pub(crate) async fn models() -> (Status, (ContentType, Json<serde_json::Value>))
let which_model = guard.which_model;
let model_obj = ModelObject {
id: which_model_to_id(which_model).to_string(),
id: which_model.openai_model_id().to_string(),
object: "model".to_string(),
created: None, // We don't track creation time
owned_by: which_model_to_owner(which_model).to_string(),
owned_by: which_model.owner().to_string(),
};
drop(guard);
@@ -373,17 +502,63 @@ mod tests {
assert_eq!(error, Some("model not initialized"));
}
#[test]
fn test_parse_embedding_input_string() {
let input = serde_json::json!("hello");
let out = parse_embedding_input(&input).unwrap();
assert_eq!(out, vec!["hello".to_string()]);
}
#[test]
fn test_parse_embedding_input_array() {
let input = serde_json::json!(["a", "b"]);
let out = parse_embedding_input(&input).unwrap();
assert_eq!(out, vec!["a".to_string(), "b".to_string()]);
}
#[test]
fn test_validate_rerank_input() {
let query = "hello";
let docs = vec!["doc1".to_string(), "doc2".to_string()];
assert!(validate_rerank_input(query, &docs).is_ok());
}
#[test]
fn test_validate_rerank_input_empty_doc() {
let query = "hello";
let docs = vec!["".to_string()];
assert!(validate_rerank_input(query, &docs).is_err());
}
// Test model type classification
#[test]
fn test_get_model_type_llm() {
assert_eq!(WhichModel::Qwen3_0_6B.model_type(), "llm");
assert_eq!(WhichModel::Qwen3vl2B.model_type(), "llm");
assert_eq!(WhichModel::MiniCPM4_0_5B.model_type(), "llm");
assert_eq!(WhichModel::Qwen2_5vl3B.model_type(), "llm");
assert_eq!(WhichModel::Qwen2_5vl7B.model_type(), "llm");
assert_eq!(WhichModel::Qwen3vl4B.model_type(), "llm");
assert_eq!(WhichModel::Qwen3vl8B.model_type(), "llm");
assert_eq!(WhichModel::Qwen3vl32B.model_type(), "llm");
}
#[test]
fn test_get_model_type_vlm() {
assert_eq!(WhichModel::Qwen3vl2B.model_type(), "vlm");
assert_eq!(WhichModel::Qwen2_5vl3B.model_type(), "vlm");
assert_eq!(WhichModel::Qwen2_5vl7B.model_type(), "vlm");
assert_eq!(WhichModel::Qwen3vl4B.model_type(), "vlm");
assert_eq!(WhichModel::Qwen3vl8B.model_type(), "vlm");
assert_eq!(WhichModel::Qwen3vl32B.model_type(), "vlm");
}
#[test]
fn test_get_model_type_embedding() {
assert_eq!(WhichModel::Qwen3Embedding0_6B.model_type(), "embedding");
assert_eq!(WhichModel::Qwen3Embedding4B.model_type(), "embedding");
assert_eq!(WhichModel::Qwen3Embedding8B.model_type(), "embedding");
}
#[test]
fn test_get_model_type_reranker() {
assert_eq!(WhichModel::Qwen3Reranker0_6B.model_type(), "reranker");
assert_eq!(WhichModel::Qwen3Reranker4B.model_type(), "reranker");
assert_eq!(WhichModel::Qwen3Reranker8B.model_type(), "reranker");
}
#[test]
@@ -412,6 +587,14 @@ mod tests {
#[test]
fn test_get_model_id() {
assert_eq!(WhichModel::Qwen3_0_6B.model_id(), "Qwen/Qwen3-0.6B");
assert_eq!(
WhichModel::Qwen3Reranker4B.model_id(),
"Qwen/Qwen3-Reranker-4B"
);
assert_eq!(
WhichModel::Qwen3Reranker8B.model_id(),
"Qwen/Qwen3-Reranker-8B"
);
assert_eq!(
WhichModel::DeepSeekOCR.model_id(),
"deepseek-ai/DeepSeek-OCR"
@@ -421,26 +604,30 @@ mod tests {
// Test OpenAI-compatible model ID conversion
#[test]
fn test_which_model_to_id() {
assert_eq!(which_model_to_id(WhichModel::Qwen3_0_6B), "qwen3-0.6b");
assert_eq!(which_model_to_id(WhichModel::DeepSeekOCR), "deepseek-ocr");
assert_eq!(which_model_to_id(WhichModel::VoxCPM1_5), "voxcpm1.5");
fn test_openai_model_id() {
assert_eq!(WhichModel::Qwen3_0_6B.openai_model_id(), "qwen3-0.6b");
assert_eq!(
which_model_to_id(WhichModel::MiniCPM4_0_5B),
"minicpm4-0.5b"
WhichModel::Qwen3Reranker4B.openai_model_id(),
"qwen3-reranker-4b"
);
assert_eq!(
WhichModel::Qwen3Reranker8B.openai_model_id(),
"qwen3-reranker-8b"
);
assert_eq!(WhichModel::DeepSeekOCR.openai_model_id(), "deepseek-ocr");
assert_eq!(WhichModel::VoxCPM1_5.openai_model_id(), "voxcpm1.5");
assert_eq!(WhichModel::MiniCPM4_0_5B.openai_model_id(), "minicpm4-0.5b");
}
// Test owner/organization mapping
#[test]
fn test_which_model_to_owner() {
assert_eq!(which_model_to_owner(WhichModel::Qwen3_0_6B), "Qwen");
assert_eq!(which_model_to_owner(WhichModel::DeepSeekOCR), "deepseek-ai");
assert_eq!(which_model_to_owner(WhichModel::VoxCPM1_5), "OpenBMB");
assert_eq!(
which_model_to_owner(WhichModel::HunyuanOCR),
"Tencent-Hunyuan"
);
fn test_model_owner() {
assert_eq!(WhichModel::Qwen3_0_6B.owner(), "Qwen");
assert_eq!(WhichModel::Qwen3Reranker4B.owner(), "Qwen");
assert_eq!(WhichModel::Qwen3Reranker8B.owner(), "Qwen");
assert_eq!(WhichModel::DeepSeekOCR.owner(), "deepseek-ai");
assert_eq!(WhichModel::VoxCPM1_5.owner(), "OpenBMB");
assert_eq!(WhichModel::HunyuanOCR.owner(), "Tencent-Hunyuan");
}
}
+3 -1
View File
@@ -27,7 +27,9 @@ impl ExecModel for FunASRNanoExec {
println!("Time elapsed in load model is: {:?}", i_duration);
// Create ChatCompletionParameters for ASR
let url = &input[1];
let url = input.get(1).ok_or_else(|| {
anyhow::anyhow!("fun-asr-nano requires a second input: audio path or URL")
})?;
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
+3 -1
View File
@@ -28,7 +28,9 @@ impl ExecModel for GlmASRNanoExec {
// Create ChatCompletionParameters for ASR
// Input should be an audio file path
let url = &input[1];
let url = input.get(1).ok_or_else(|| {
anyhow::anyhow!("glm-asr-nano requires a second input: audio path or URL")
})?;
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
+2
View File
@@ -14,6 +14,8 @@ pub mod qwen2_5vl;
pub mod qwen3;
pub mod qwen3_5;
pub mod qwen3_asr;
pub mod qwen3_embedding;
pub mod qwen3_reranker;
pub mod qwen3vl;
pub mod rmbg2_0;
pub mod voxcpm;
+3 -1
View File
@@ -19,7 +19,9 @@ impl ExecModel for Qwen2_5vlExec {
} else {
input_text.clone()
};
let url = &input[1];
let url = input.get(1).ok_or_else(|| {
anyhow::anyhow!("qwen2.5vl requires a second input: image path or URL")
})?;
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
+3 -1
View File
@@ -148,7 +148,9 @@ impl ExecModel for Qwen3_5Exec {
let mut model = Qwen3_5GenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let url = &input[1];
let url = input
.get(1)
.ok_or_else(|| anyhow!("qwen3.5 requires a second input: image/video path or URL"))?;
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
+39
View File
@@ -0,0 +1,39 @@
use std::time::Instant;
use anyhow::Result;
use crate::exec::ExecModel;
use crate::models::qwen3_embedding::generate::Qwen3EmbeddingModel;
use crate::utils::get_file_path;
pub struct Qwen3EmbeddingExec;
impl ExecModel for Qwen3EmbeddingExec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
let input_text = input
.first()
.ok_or_else(|| anyhow::anyhow!("embedding run requires one text input"))?;
let text = if input_text.starts_with("file://") {
let path = get_file_path(input_text)?;
std::fs::read_to_string(path)?
} else {
input_text.clone()
};
let i_start = Instant::now();
let mut model = Qwen3EmbeddingModel::init(weight_path, None, None)?;
println!("Time elapsed in load model is: {:?}", i_start.elapsed());
let i_start = Instant::now();
let embedding = model.embed(&[text])?;
println!("Time elapsed in embedding is: {:?}", i_start.elapsed());
let output_json = serde_json::to_string_pretty(&embedding)?;
println!("{}", output_json);
if let Some(out) = output {
std::fs::write(out, output_json)?;
println!("Output saved to: {}", out);
}
Ok(())
}
}
+87
View File
@@ -0,0 +1,87 @@
use std::time::Instant;
use anyhow::{Result, anyhow};
use serde::Serialize;
use crate::exec::ExecModel;
use crate::models::qwen3_reranker::generate::Qwen3RerankerModel;
#[derive(Debug, Serialize)]
struct RerankItem {
index: usize,
score: f32,
document: String,
}
pub struct Qwen3RerankerExec;
impl ExecModel for Qwen3RerankerExec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
if input.len() < 2 {
return Err(anyhow!(
"reranker run requires two inputs: <query> <documents-source>"
));
}
let query = input[0].clone();
let docs_source = input[1].clone();
let documents = parse_documents_source(&docs_source)?;
if documents.is_empty() {
return Err(anyhow!("documents list is empty"));
}
let mut model = Qwen3RerankerModel::init(weight_path, None, None)?;
let i_start = Instant::now();
let scores = model.rerank(&query, &documents)?;
println!("Time elapsed in rerank is: {:?}", i_start.elapsed());
let mut ranked = scores
.into_iter()
.enumerate()
.map(|(index, score)| RerankItem {
index,
score,
document: documents[index].clone(),
})
.collect::<Vec<_>>();
ranked.sort_by(|a, b| b.score.total_cmp(&a.score));
let output_json = serde_json::to_string_pretty(&ranked)?;
println!("{}", output_json);
let output_path = output
.map(|o| o.to_string())
.unwrap_or_else(|| format!("qwen3-rerank-{}.json", chrono::Utc::now().timestamp()));
std::fs::write(&output_path, output_json.as_bytes())?;
println!("Generate rerank output to {}", output_path);
Ok(())
}
}
fn parse_documents_source(source: &str) -> Result<Vec<String>> {
if source.starts_with("file://") {
let path = source.trim_start_matches("file://");
return read_documents_file(path);
}
if std::path::Path::new(source).exists() {
return read_documents_file(source);
}
let docs = source
.split("|||")
.map(str::trim)
.filter(|x| !x.is_empty())
.map(|x| x.to_string())
.collect::<Vec<_>>();
Ok(docs)
}
fn read_documents_file(path: &str) -> Result<Vec<String>> {
let content = std::fs::read_to_string(path)?;
let docs = content
.lines()
.map(str::trim)
.filter(|line| !line.is_empty())
.map(|line| line.to_string())
.collect::<Vec<_>>();
Ok(docs)
}
+101 -68
View File
@@ -2,7 +2,7 @@ use std::sync::atomic::{AtomicBool, Ordering};
use std::{net::IpAddr, str::FromStr, sync::Arc};
use aha::{
models::WhichModel,
models::{LISTED_MODELS, ModelArtifactFormat, WhichModel},
process::{cleanup_pid_file, create_pid_file},
utils::{download_model, get_default_save_dir},
};
@@ -242,37 +242,9 @@ struct ModelInfo {
/// List all supported models
fn run_list(args: ListArgs) -> anyhow::Result<()> {
let models = [
WhichModel::MiniCPM4_0_5B,
WhichModel::Qwen2_5vl3B,
WhichModel::Qwen2_5vl7B,
WhichModel::Qwen3_0_6B,
WhichModel::Qwen3_5_0_8B,
WhichModel::Qwen3_5_2B,
WhichModel::Qwen3_5_4B,
WhichModel::Qwen3_5_9B,
WhichModel::Qwen3ASR0_6B,
WhichModel::Qwen3ASR1_7B,
WhichModel::Qwen3vl2B,
WhichModel::Qwen3vl4B,
WhichModel::Qwen3vl8B,
WhichModel::Qwen3vl32B,
WhichModel::DeepSeekOCR,
WhichModel::DeepSeekOCR2,
WhichModel::HunyuanOCR,
WhichModel::PaddleOCRVL,
WhichModel::PaddleOCRVL1_5,
WhichModel::RMBG2_0,
WhichModel::VoxCPM,
WhichModel::VoxCPM1_5,
WhichModel::GlmASRNano2512,
WhichModel::FunASRNano2512,
WhichModel::GlmOCR,
];
if args.json {
// JSON output
let model_infos: Vec<ModelInfo> = models
let model_infos: Vec<ModelInfo> = LISTED_MODELS
.iter()
.map(|model| {
let possible_value = model.to_possible_value().unwrap();
@@ -294,11 +266,11 @@ fn run_list(args: ListArgs) -> anyhow::Result<()> {
"Model Name", "ModelScope ID", "Download"
);
println!("{}", "-".repeat(80));
for model in models {
for model in LISTED_MODELS {
let possible_value = model.to_possible_value().unwrap();
let name = possible_value.get_name();
let id = model.model_id();
let download_status = if is_model_downloaded(model) {
let download_status = if is_model_downloaded(*model) {
""
} else {
""
@@ -310,6 +282,46 @@ fn run_list(args: ListArgs) -> anyhow::Result<()> {
Ok(())
}
async fn resolve_model_paths_for_server(
model: WhichModel,
weight_path: Option<String>,
save_dir: Option<String>,
download_retries: Option<u32>,
gguf_path: Option<String>,
mmproj_path: Option<String>,
allow_download: bool,
) -> anyhow::Result<(String, Option<String>, Option<String>)> {
match model.artifact_format() {
ModelArtifactFormat::Gguf => {
if gguf_path.is_none() {
return Err(anyhow!("gguf model path is required"));
}
Ok(("GGUF".to_string(), gguf_path, mmproj_path))
}
ModelArtifactFormat::Safetensors => {
let model_id = model.model_id();
let model_path = match weight_path {
Some(path) => path,
None if allow_download => {
let save_dir = match save_dir {
Some(dir) => dir,
None => get_default_save_dir().expect("Failed to get home directory"),
};
let max_retries = download_retries.unwrap_or(3);
download_model(model_id, &save_dir, max_retries).await?;
save_dir + "/" + model_id
}
None => get_default_weight_path(model),
};
Ok((model_path, None, None))
}
ModelArtifactFormat::Onnx => Err(anyhow!(
"onnx runtime is not integrated yet for model {}",
model.openai_model_id()
)),
}
}
/// Run the 'cli' subcommand: download model (if needed) and start service
async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
let CliArgs {
@@ -320,28 +332,16 @@ async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
gguf_path,
mmproj_path,
} = args;
let model_id = common.model.model_id();
let (model_path, gguf, mmproj) = if model_id.eq("GGUF") {
if gguf_path.is_none() {
return Err(anyhow!("gguf model path is required"));
}
("GGUF".to_string(), gguf_path, mmproj_path)
} else {
let model_path = match weight_path {
Some(path) => path,
None => {
let save_dir = match save_dir {
Some(dir) => dir,
None => get_default_save_dir().expect("Failed to get home directory"),
};
let max_retries = download_retries.unwrap_or(3);
download_model(model_id, &save_dir, max_retries).await?;
save_dir + "/" + model_id
}
};
(model_path, None, None)
};
let (model_path, gguf, mmproj) = resolve_model_paths_for_server(
common.model,
weight_path,
save_dir,
download_retries,
gguf_path,
mmproj_path,
true,
)
.await?;
init(common.model, model_path, gguf, mmproj)?;
start_http_server(common.address, common.port, common.allow_remote_shutdown).await?;
@@ -357,19 +357,16 @@ async fn run_serv(args: ServArgs) -> anyhow::Result<()> {
gguf_path,
mmproj_path,
} = args;
let model_id = common.model.model_id();
let (model_path, gguf, mmproj) = if model_id.eq("GGUF") {
if gguf_path.is_none() {
return Err(anyhow!("gguf model path is required"));
}
("GGUF".to_string(), gguf_path, mmproj_path)
} else {
let model_path = match weight_path {
Some(path) => path,
None => get_default_weight_path(common.model),
};
(model_path, None, None)
};
let (model_path, gguf, mmproj) = resolve_model_paths_for_server(
common.model,
weight_path,
None,
None,
gguf_path,
mmproj_path,
false,
)
.await?;
init(common.model, model_path, gguf, mmproj)?;
start_http_server(common.address, common.port, common.allow_remote_shutdown).await?;
@@ -426,6 +423,12 @@ async fn run_download(args: DownloadArgs) -> anyhow::Result<()> {
download_retries,
} = args;
let model_id = model.model_id();
if !model.is_download_managed() {
return Err(anyhow!(
"{} does not use managed model download. Please provide local artifact path directly when serving/running.",
model.openai_model_id()
));
}
let save_dir = match save_dir {
Some(dir) => dir,
@@ -473,6 +476,18 @@ fn run_run(args: RunArgs) -> anyhow::Result<()> {
use aha::exec::qwen3::Qwen3Exec;
Qwen3Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B
| WhichModel::Qwen3Embedding8B => {
use aha::exec::qwen3_embedding::Qwen3EmbeddingExec;
Qwen3EmbeddingExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3Reranker0_6B
| WhichModel::Qwen3Reranker4B
| WhichModel::Qwen3Reranker8B => {
use aha::exec::qwen3_reranker::Qwen3RerankerExec;
Qwen3RerankerExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3_5_0_8B => {
use aha::exec::qwen3_5::Qwen3_5Exec;
Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?;
@@ -489,7 +504,19 @@ fn run_run(args: RunArgs) -> anyhow::Result<()> {
use aha::exec::qwen3_5::Qwen3_5Exec;
Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3_5Gguf => {
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2 => {
use aha::exec::qwen3_5::Qwen3_5Exec;
Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3_5Gguf
| WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_0_8BUnslothGguf
| WhichModel::Qwen3_5_2BUnslothGguf
| WhichModel::Qwen3_5_4BUnslothGguf
| WhichModel::Qwen3_5_0_8BLmstudioGguf
| WhichModel::Qwen3_5_2BLmstudioGguf
| WhichModel::Qwen3_5_4BLmstudioGguf => {
use aha::exec::qwen3_5::Qwen3_5Exec;
Qwen3_5Exec::run_gguf(&input, output.as_deref(), gguf_path, mmproj_path)?;
}
@@ -734,6 +761,12 @@ pub(crate) async fn start_http_server(
builder = builder.mount("/audio", routes![api::speech, api::transcriptions]);
// /v1/audio/transcriptions (OpenAI standard ASR transcription endpoint)
builder = builder.mount("/v1/audio", routes![api::transcriptions]);
// /embeddings and /v1/embeddings (OpenAI-compatible embeddings endpoint)
builder = builder.mount("/", routes![api::embeddings]);
builder = builder.mount("/v1", routes![api::embeddings]);
// /rerank and /v1/rerank
builder = builder.mount("/", routes![api::rerank]);
builder = builder.mount("/v1", routes![api::rerank]);
// Health check and model info endpoints
builder = builder.mount("/", routes![api::health, api::models]);
// Shutdown endpoint
+1
View File
@@ -8,6 +8,7 @@ use candle_nn::{
};
pub mod gguf;
pub mod retrieval;
use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
+41
View File
@@ -0,0 +1,41 @@
use anyhow::{Result, anyhow};
pub trait TextEmbeddingBackend {
fn embed_texts(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>>;
}
pub fn l2_normalize(v: &mut [f32]) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in v.iter_mut() {
*x /= norm;
}
}
}
pub fn mean_pool(embeddings: &[Vec<f32>]) -> Result<Vec<f32>> {
let first = embeddings
.first()
.ok_or_else(|| anyhow!("embedding hidden state is empty"))?;
let mut pooled = vec![0f32; first.len()];
for row in embeddings {
if row.len() != first.len() {
return Err(anyhow!("inconsistent embedding width in hidden state"));
}
for (idx, value) in row.iter().enumerate() {
pooled[idx] += *value;
}
}
let inv = 1.0f32 / embeddings.len() as f32;
for value in &mut pooled {
*value *= inv;
}
Ok(pooled)
}
pub fn cosine_similarity(lhs: &[f32], rhs: &[f32]) -> Result<f32> {
if lhs.len() != rhs.len() {
return Err(anyhow!("embedding dimension mismatch"));
}
Ok(lhs.iter().zip(rhs.iter()).map(|(l, r)| l * r).sum::<f32>())
}
+299 -2
View File
@@ -15,6 +15,8 @@ pub mod qwen2_5vl;
pub mod qwen3;
pub mod qwen3_5;
pub mod qwen3_asr;
pub mod qwen3_embedding;
pub mod qwen3_reranker;
pub mod qwen3vl;
pub mod rmbg2_0;
pub mod voxcpm;
@@ -33,10 +35,18 @@ use crate::models::{
hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel,
paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel,
qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel,
qwen3_asr::generate::Qwen3AsrGenerateModel, qwen3vl::generate::Qwen3VLGenerateModel,
qwen3_asr::generate::Qwen3AsrGenerateModel, qwen3_embedding::generate::Qwen3EmbeddingModel,
qwen3_reranker::generate::Qwen3RerankerModel, qwen3vl::generate::Qwen3VLGenerateModel,
rmbg2_0::generate::RMBG2_0Model, voxcpm::generate::VoxCPMGenerate,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelArtifactFormat {
Safetensors,
Gguf,
Onnx,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
pub enum WhichModel {
#[value(name = "minicpm4-0.5b", hide = true)]
@@ -47,6 +57,18 @@ pub enum WhichModel {
Qwen2_5vl7B,
#[value(name = "qwen3-0.6b", hide = true)]
Qwen3_0_6B,
#[value(name = "qwen3-embedding-0.6b", hide = true)]
Qwen3Embedding0_6B,
#[value(name = "qwen3-embedding-4b", hide = true)]
Qwen3Embedding4B,
#[value(name = "qwen3-embedding-8b", hide = true)]
Qwen3Embedding8B,
#[value(name = "qwen3-reranker-0.6b", hide = true)]
Qwen3Reranker0_6B,
#[value(name = "qwen3-reranker-4b", hide = true)]
Qwen3Reranker4B,
#[value(name = "qwen3-reranker-8b", hide = true)]
Qwen3Reranker8B,
#[value(name = "qwen3.5-0.8b", hide = true)]
Qwen3_5_0_8B,
#[value(name = "qwen3.5-2b", hide = true)]
@@ -55,8 +77,35 @@ pub enum WhichModel {
Qwen3_5_4B,
#[value(name = "qwen3.5-9b", hide = true)]
Qwen3_5_9B,
#[value(
name = "qwen3.5-9b-claude-4.6-opus-reasoning-distilled-v2",
hide = true
)]
Qwen3_5_9BClaude46OpusReasoningDistilledV2,
#[value(name = "qwen3.5-gguf", hide = true)]
Qwen3_5Gguf,
#[value(
name = "qwen3.5-4b-claude-4.6-opus-reasoning-distilled-v2-gguf",
hide = true
)]
Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf,
#[value(
name = "qwen3.5-9b-claude-4.6-opus-reasoning-distilled-v2-gguf",
hide = true
)]
Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf,
#[value(name = "qwen3.5-0.8b-unsloth-gguf", hide = true)]
Qwen3_5_0_8BUnslothGguf,
#[value(name = "qwen3.5-2b-unsloth-gguf", hide = true)]
Qwen3_5_2BUnslothGguf,
#[value(name = "qwen3.5-4b-unsloth-gguf", hide = true)]
Qwen3_5_4BUnslothGguf,
#[value(name = "qwen3.5-0.8b-lmstudio-gguf", hide = true)]
Qwen3_5_0_8BLmstudioGguf,
#[value(name = "qwen3.5-2b-lmstudio-gguf", hide = true)]
Qwen3_5_2BLmstudioGguf,
#[value(name = "qwen3.5-4b-lmstudio-gguf", hide = true)]
Qwen3_5_4BLmstudioGguf,
#[value(name = "qwen3asr-0.6b", hide = true)]
Qwen3ASR0_6B,
#[value(name = "qwen3asr-1.7b", hide = true)]
@@ -93,7 +142,165 @@ pub enum WhichModel {
GlmOCR,
}
pub const LISTED_MODELS: &[WhichModel] = &[
WhichModel::MiniCPM4_0_5B,
WhichModel::Qwen2_5vl3B,
WhichModel::Qwen2_5vl7B,
WhichModel::Qwen3_0_6B,
WhichModel::Qwen3Embedding0_6B,
WhichModel::Qwen3Embedding4B,
WhichModel::Qwen3Embedding8B,
WhichModel::Qwen3Reranker0_6B,
WhichModel::Qwen3Reranker4B,
WhichModel::Qwen3Reranker8B,
WhichModel::Qwen3_5_0_8B,
WhichModel::Qwen3_5_2B,
WhichModel::Qwen3_5_4B,
WhichModel::Qwen3_5_9B,
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2,
WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf,
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf,
WhichModel::Qwen3_5_0_8BUnslothGguf,
WhichModel::Qwen3_5_2BUnslothGguf,
WhichModel::Qwen3_5_4BUnslothGguf,
WhichModel::Qwen3_5_0_8BLmstudioGguf,
WhichModel::Qwen3_5_2BLmstudioGguf,
WhichModel::Qwen3_5_4BLmstudioGguf,
WhichModel::Qwen3ASR0_6B,
WhichModel::Qwen3ASR1_7B,
WhichModel::Qwen3vl2B,
WhichModel::Qwen3vl4B,
WhichModel::Qwen3vl8B,
WhichModel::Qwen3vl32B,
WhichModel::DeepSeekOCR,
WhichModel::DeepSeekOCR2,
WhichModel::HunyuanOCR,
WhichModel::PaddleOCRVL,
WhichModel::PaddleOCRVL1_5,
WhichModel::RMBG2_0,
WhichModel::VoxCPM,
WhichModel::VoxCPM1_5,
WhichModel::GlmASRNano2512,
WhichModel::FunASRNano2512,
WhichModel::GlmOCR,
];
impl WhichModel {
pub fn artifact_format(self) -> ModelArtifactFormat {
match self {
WhichModel::Qwen3_5Gguf
| WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_0_8BUnslothGguf
| WhichModel::Qwen3_5_2BUnslothGguf
| WhichModel::Qwen3_5_4BUnslothGguf
| WhichModel::Qwen3_5_0_8BLmstudioGguf
| WhichModel::Qwen3_5_2BLmstudioGguf
| WhichModel::Qwen3_5_4BLmstudioGguf => ModelArtifactFormat::Gguf,
_ => ModelArtifactFormat::Safetensors,
}
}
pub fn is_download_managed(self) -> bool {
!matches!(
self.artifact_format(),
ModelArtifactFormat::Gguf | ModelArtifactFormat::Onnx
)
}
pub fn openai_model_id(self) -> &'static str {
match self {
WhichModel::MiniCPM4_0_5B => "minicpm4-0.5b",
WhichModel::Qwen2_5vl3B => "qwen2.5vl-3b",
WhichModel::Qwen2_5vl7B => "qwen2.5vl-7b",
WhichModel::Qwen3_0_6B => "qwen3-0.6b",
WhichModel::Qwen3Embedding0_6B => "qwen3-embedding-0.6b",
WhichModel::Qwen3Embedding4B => "qwen3-embedding-4b",
WhichModel::Qwen3Embedding8B => "qwen3-embedding-8b",
WhichModel::Qwen3Reranker0_6B => "qwen3-reranker-0.6b",
WhichModel::Qwen3Reranker4B => "qwen3-reranker-4b",
WhichModel::Qwen3Reranker8B => "qwen3-reranker-8b",
WhichModel::Qwen3_5_0_8B => "qwen3.5-0.8b",
WhichModel::Qwen3_5_2B => "qwen3.5-2b",
WhichModel::Qwen3_5_4B => "qwen3.5-4b",
WhichModel::Qwen3_5_9B => "qwen3.5-9b",
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2 => {
"qwen3.5-9b-claude-4.6-opus-reasoning-distilled-v2"
}
WhichModel::Qwen3_5Gguf => "qwen3.5-gguf",
WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf => {
"qwen3.5-4b-claude-4.6-opus-reasoning-distilled-v2-gguf"
}
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf => {
"qwen3.5-9b-claude-4.6-opus-reasoning-distilled-v2-gguf"
}
WhichModel::Qwen3_5_0_8BUnslothGguf => "qwen3.5-0.8b-unsloth-gguf",
WhichModel::Qwen3_5_2BUnslothGguf => "qwen3.5-2b-unsloth-gguf",
WhichModel::Qwen3_5_4BUnslothGguf => "qwen3.5-4b-unsloth-gguf",
WhichModel::Qwen3_5_0_8BLmstudioGguf => "qwen3.5-0.8b-lmstudio-gguf",
WhichModel::Qwen3_5_2BLmstudioGguf => "qwen3.5-2b-lmstudio-gguf",
WhichModel::Qwen3_5_4BLmstudioGguf => "qwen3.5-4b-lmstudio-gguf",
WhichModel::Qwen3ASR0_6B => "qwen3asr-0.6b",
WhichModel::Qwen3ASR1_7B => "qwen3asr-1.7b",
WhichModel::Qwen3vl2B => "qwen3vl-2b",
WhichModel::Qwen3vl4B => "qwen3vl-4b",
WhichModel::Qwen3vl8B => "qwen3vl-8b",
WhichModel::Qwen3vl32B => "qwen3vl-32b",
WhichModel::DeepSeekOCR => "deepseek-ocr",
WhichModel::DeepSeekOCR2 => "deepseek-ocr2",
WhichModel::HunyuanOCR => "hunyuan-ocr",
WhichModel::PaddleOCRVL => "paddleocr-vl",
WhichModel::PaddleOCRVL1_5 => "paddleocr-vl1.5",
WhichModel::RMBG2_0 => "rmbg2.0",
WhichModel::VoxCPM => "voxcpm",
WhichModel::VoxCPM1_5 => "voxcpm1.5",
WhichModel::GlmASRNano2512 => "glm-asr-nano-2512",
WhichModel::FunASRNano2512 => "fun-asr-nano-2512",
WhichModel::GlmOCR => "glm-ocr",
}
}
pub fn owner(self) -> &'static str {
match self {
WhichModel::MiniCPM4_0_5B => "OpenBMB",
WhichModel::Qwen2_5vl3B | WhichModel::Qwen2_5vl7B => "Qwen",
WhichModel::Qwen3_0_6B
| WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B
| WhichModel::Qwen3Embedding8B
| WhichModel::Qwen3Reranker0_6B
| WhichModel::Qwen3Reranker4B
| WhichModel::Qwen3Reranker8B
| WhichModel::Qwen3ASR0_6B
| WhichModel::Qwen3ASR1_7B => "Qwen",
WhichModel::Qwen3vl2B
| WhichModel::Qwen3vl4B
| WhichModel::Qwen3vl8B
| WhichModel::Qwen3vl32B
| WhichModel::Qwen3_5Gguf => "Qwen",
WhichModel::Qwen3_5_0_8B
| WhichModel::Qwen3_5_2B
| WhichModel::Qwen3_5_4B
| WhichModel::Qwen3_5_9B => "Qwen",
WhichModel::Qwen3_5_0_8BUnslothGguf
| WhichModel::Qwen3_5_2BUnslothGguf
| WhichModel::Qwen3_5_4BUnslothGguf => "unsloth",
WhichModel::Qwen3_5_0_8BLmstudioGguf
| WhichModel::Qwen3_5_2BLmstudioGguf
| WhichModel::Qwen3_5_4BLmstudioGguf => "lmstudio-community",
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2
| WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf => "Jackrong",
WhichModel::DeepSeekOCR | WhichModel::DeepSeekOCR2 => "deepseek-ai",
WhichModel::HunyuanOCR => "Tencent-Hunyuan",
WhichModel::PaddleOCRVL | WhichModel::PaddleOCRVL1_5 => "PaddlePaddle",
WhichModel::RMBG2_0 => "AI-ModelScope",
WhichModel::VoxCPM | WhichModel::VoxCPM1_5 => "OpenBMB",
WhichModel::GlmASRNano2512 | WhichModel::GlmOCR => "ZhipuAI",
WhichModel::FunASRNano2512 => "FunAudioLLM",
}
}
/// Get the ModelScope model ID for this model variant
pub fn model_id(self) -> &'static str {
match self {
@@ -101,11 +308,32 @@ impl WhichModel {
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
WhichModel::Qwen3_0_6B => "Qwen/Qwen3-0.6B",
WhichModel::Qwen3Embedding0_6B => "Qwen/Qwen3-Embedding-0.6B",
WhichModel::Qwen3Embedding4B => "Qwen/Qwen3-Embedding-4B",
WhichModel::Qwen3Embedding8B => "Qwen/Qwen3-Embedding-8B",
WhichModel::Qwen3Reranker0_6B => "Qwen/Qwen3-Reranker-0.6B",
WhichModel::Qwen3Reranker4B => "Qwen/Qwen3-Reranker-4B",
WhichModel::Qwen3Reranker8B => "Qwen/Qwen3-Reranker-8B",
WhichModel::Qwen3_5_0_8B => "Qwen/Qwen3.5-0.8B",
WhichModel::Qwen3_5_2B => "Qwen/Qwen3.5-2B",
WhichModel::Qwen3_5_4B => "Qwen/Qwen3.5-4B",
WhichModel::Qwen3_5_9B => "Qwen/Qwen3.5-9B",
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2 => {
"Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2"
}
WhichModel::Qwen3_5Gguf => "GGUF",
WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf => {
"Jackrong/Qwen3.5-4B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF"
}
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf => {
"Jackrong/Qwen3.5-9B-Claude-4.6-Opus-Reasoning-Distilled-v2-GGUF"
}
WhichModel::Qwen3_5_0_8BUnslothGguf => "unsloth/Qwen3.5-0.8B-GGUF",
WhichModel::Qwen3_5_2BUnslothGguf => "unsloth/Qwen3.5-2B-GGUF",
WhichModel::Qwen3_5_4BUnslothGguf => "unsloth/Qwen3.5-4B-GGUF",
WhichModel::Qwen3_5_0_8BLmstudioGguf => "lmstudio-community/Qwen3.5-0.8B-GGUF",
WhichModel::Qwen3_5_2BLmstudioGguf => "lmstudio-community/Qwen3.5-2B-GGUF",
WhichModel::Qwen3_5_4BLmstudioGguf => "lmstudio-community/Qwen3.5-4B-GGUF",
WhichModel::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B",
WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B",
WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct",
@@ -131,6 +359,12 @@ impl WhichModel {
match self {
// LLM models
WhichModel::MiniCPM4_0_5B | WhichModel::Qwen3_0_6B => "llm",
WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B
| WhichModel::Qwen3Embedding8B => "embedding",
WhichModel::Qwen3Reranker0_6B
| WhichModel::Qwen3Reranker4B
| WhichModel::Qwen3Reranker8B => "reranker",
WhichModel::Qwen2_5vl3B
| WhichModel::Qwen2_5vl7B
| WhichModel::Qwen3vl2B
@@ -141,7 +375,16 @@ impl WhichModel {
| WhichModel::Qwen3_5_2B
| WhichModel::Qwen3_5_4B
| WhichModel::Qwen3_5_9B
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2
| WhichModel::Qwen3_5Gguf => "vlm",
WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_0_8BUnslothGguf
| WhichModel::Qwen3_5_2BUnslothGguf
| WhichModel::Qwen3_5_4BUnslothGguf
| WhichModel::Qwen3_5_0_8BLmstudioGguf
| WhichModel::Qwen3_5_2BLmstudioGguf
| WhichModel::Qwen3_5_4BLmstudioGguf => "vlm",
// OCR models
WhichModel::DeepSeekOCR
| WhichModel::DeepSeekOCR2
@@ -179,6 +422,8 @@ pub enum ModelInstance<'a> {
MiniCPM4(MiniCPMGenerateModel<'a>),
Qwen2_5VL(Qwen2_5VLGenerateModel<'a>),
Qwen3(Qwen3GenerateModel<'a>),
Qwen3Embedding(Qwen3EmbeddingModel),
Qwen3Reranker(Qwen3RerankerModel),
Qwen3_5(Qwen3_5GenerateModel<'a>),
Qwen3ASR(Qwen3AsrGenerateModel<'a>),
Qwen3VL(Box<Qwen3VLGenerateModel<'a>>),
@@ -198,6 +443,12 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::MiniCPM4(model) => model.generate(mes),
ModelInstance::Qwen2_5VL(model) => model.generate(mes),
ModelInstance::Qwen3(model) => model.generate(mes),
ModelInstance::Qwen3Embedding(_) => {
Err(anyhow!("embedding model does not support chat completions"))
}
ModelInstance::Qwen3Reranker(_) => {
Err(anyhow!("reranker model does not support chat completions"))
}
ModelInstance::Qwen3_5(model) => model.generate(mes),
ModelInstance::Qwen3ASR(model) => model.generate(mes),
ModelInstance::Qwen3VL(model) => model.generate(mes),
@@ -227,6 +478,12 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::MiniCPM4(model) => model.generate_stream(mes),
ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes),
ModelInstance::Qwen3(model) => model.generate_stream(mes),
ModelInstance::Qwen3Embedding(_) => Err(anyhow!(
"embedding model does not support streaming chat completions"
)),
ModelInstance::Qwen3Reranker(_) => Err(anyhow!(
"reranker model does not support streaming chat completions"
)),
ModelInstance::Qwen3_5(model) => model.generate_stream(mes),
ModelInstance::Qwen3VL(model) => model.generate_stream(mes),
ModelInstance::Qwen3ASR(model) => model.generate_stream(mes),
@@ -242,6 +499,22 @@ impl<'a> GenerateModel for ModelInstance<'a> {
}
}
impl<'a> ModelInstance<'a> {
pub fn embedding(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>> {
match self {
ModelInstance::Qwen3Embedding(model) => model.embed(input),
_ => Err(anyhow!("current model does not support embeddings")),
}
}
pub fn rerank(&mut self, query: &str, documents: &[String]) -> Result<Vec<f32>> {
match self {
ModelInstance::Qwen3Reranker(model) => model.rerank(query, documents),
_ => Err(anyhow!("current model does not support reranking")),
}
}
}
pub fn load_model<'a>(
model_type: WhichModel,
path: &str,
@@ -265,6 +538,18 @@ pub fn load_model<'a>(
let model = Qwen3GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3(model)
}
WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B
| WhichModel::Qwen3Embedding8B => {
let model = Qwen3EmbeddingModel::init(path, None, None)?;
ModelInstance::Qwen3Embedding(model)
}
WhichModel::Qwen3Reranker0_6B
| WhichModel::Qwen3Reranker4B
| WhichModel::Qwen3Reranker8B => {
let model = Qwen3RerankerModel::init(path, None, None)?;
ModelInstance::Qwen3Reranker(model)
}
WhichModel::Qwen3_5_0_8B => {
let model = Qwen3_5GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3_5(model)
@@ -281,7 +566,19 @@ pub fn load_model<'a>(
let model = Qwen3_5GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3_5(model)
}
WhichModel::Qwen3_5Gguf => {
WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2 => {
let model = Qwen3_5GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3_5(model)
}
WhichModel::Qwen3_5Gguf
| WhichModel::Qwen3_5_4BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_9BClaude46OpusReasoningDistilledV2Gguf
| WhichModel::Qwen3_5_0_8BUnslothGguf
| WhichModel::Qwen3_5_2BUnslothGguf
| WhichModel::Qwen3_5_4BUnslothGguf
| WhichModel::Qwen3_5_0_8BLmstudioGguf
| WhichModel::Qwen3_5_2BLmstudioGguf
| WhichModel::Qwen3_5_4BLmstudioGguf => {
if gguf.is_none() {
return Err(anyhow!("Qwen3_5Gguf gguf model path is required"));
}
+19 -4
View File
@@ -205,7 +205,11 @@ pub struct Qwen3Model {
impl Qwen3Model {
pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result<Self> {
let vb = vb.pp("model");
let vb = if vb.contains_tensor("model.embed_tokens.weight") {
vb.pp("model")
} else {
vb
};
let vocab_size = config.vocab_size;
let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
@@ -235,6 +239,19 @@ impl Qwen3Model {
input_ids: Option<&Tensor>,
inputs_embeds: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
let hidden_states = self.forward_hidden(input_ids, inputs_embeds, seqlen_offset)?;
let seq_len = hidden_states.dim(1)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
}
pub fn forward_hidden(
&mut self,
input_ids: Option<&Tensor>,
inputs_embeds: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
if input_ids.is_none() && inputs_embeds.is_none() {
return Err(anyhow::anyhow!(
@@ -271,9 +288,7 @@ impl Qwen3Model {
decode_layer.forward(&hidden_states, &cos, &sin, attention_mask.as_ref())?;
}
hidden_states = self.norm.forward(&hidden_states)?;
let hidden_state = hidden_states.narrow(1, seq_len - 1, 1)?;
let logits = self.lm_head.forward(&hidden_state)?;
Ok(logits)
Ok(hidden_states)
}
pub fn embedding_token_id(&self, input_ids: &Tensor) -> Result<Tensor> {
Ok(self.embed_tokens.forward(input_ids)?)
+27
View File
@@ -0,0 +1,27 @@
use anyhow::Result;
use crate::models::qwen3::config::Qwen3Config;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Qwen3EmbeddingPoolingStrategy {
Mean,
}
#[derive(Debug, Clone)]
pub struct Qwen3EmbeddingConfig {
pub base: Qwen3Config,
pub pooling: Qwen3EmbeddingPoolingStrategy,
pub normalize: bool,
}
impl Qwen3EmbeddingConfig {
pub fn load(path: &str) -> Result<Self> {
let config_path = format!("{path}/config.json");
let base: Qwen3Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
Ok(Self {
base,
pooling: Qwen3EmbeddingPoolingStrategy::Mean,
normalize: true,
})
}
}
+27
View File
@@ -0,0 +1,27 @@
use anyhow::Result;
use candle_core::{DType, Device};
use crate::models::{
common::retrieval::TextEmbeddingBackend, qwen3_embedding::model::Qwen3EmbeddingBackend,
};
pub struct Qwen3EmbeddingModel {
backend: Qwen3EmbeddingBackend,
}
impl Qwen3EmbeddingModel {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let backend = Qwen3EmbeddingBackend::load(path, device, dtype)?;
Ok(Self { backend })
}
pub fn embed(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>> {
self.backend.embed_texts(input)
}
}
impl TextEmbeddingBackend for Qwen3EmbeddingModel {
fn embed_texts(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>> {
self.backend.embed_texts(input)
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod config;
pub mod generate;
pub mod model;
+69
View File
@@ -0,0 +1,69 @@
use anyhow::{Result, anyhow};
use candle_core::{DType, Device};
use candle_nn::VarBuilder;
use crate::{
models::{
common::retrieval::{l2_normalize, mean_pool},
qwen3::model::Qwen3Model,
qwen3_embedding::config::{Qwen3EmbeddingConfig, Qwen3EmbeddingPoolingStrategy},
},
tokenizer::TokenizerModel,
utils::{find_type_files, get_device, get_dtype},
};
pub struct Qwen3EmbeddingBackend {
tokenizer: TokenizerModel,
model: Qwen3Model,
device: Device,
pooling: Qwen3EmbeddingPoolingStrategy,
normalize: bool,
}
impl Qwen3EmbeddingBackend {
pub fn load(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let tokenizer = TokenizerModel::init(path)?;
let cfg = Qwen3EmbeddingConfig::load(path)?;
let device = get_device(device);
let dtype = get_dtype(dtype, cfg.base.torch_dtype.as_str());
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
let model = Qwen3Model::new(&cfg.base, vb)?;
Ok(Self {
tokenizer,
model,
device,
pooling: cfg.pooling,
normalize: cfg.normalize,
})
}
pub fn embed_texts(&mut self, input: &[String]) -> Result<Vec<Vec<f32>>> {
if input.is_empty() {
return Err(anyhow!("embedding input cannot be empty"));
}
let mut out = Vec::with_capacity(input.len());
for text in input {
out.push(self.embed_one(text)?);
self.model.clear_kv_cache();
}
Ok(out)
}
fn embed_one(&mut self, text: &str) -> Result<Vec<f32>> {
let input_ids = self.tokenizer.text_encode(text.to_string(), &self.device)?;
let hidden = self
.model
.forward_hidden(Some(&input_ids), None, 0)?
.squeeze(0)?
.to_dtype(DType::F32)?;
let hidden_vec = hidden.to_vec2::<f32>()?;
let mut pooled = match self.pooling {
Qwen3EmbeddingPoolingStrategy::Mean => mean_pool(&hidden_vec)?,
};
if self.normalize {
l2_normalize(&mut pooled);
}
Ok(pooled)
}
}
+17
View File
@@ -0,0 +1,17 @@
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Qwen3RerankerSimilarity {
Cosine,
}
#[derive(Debug, Clone)]
pub struct Qwen3RerankerConfig {
pub similarity: Qwen3RerankerSimilarity,
}
impl Default for Qwen3RerankerConfig {
fn default() -> Self {
Self {
similarity: Qwen3RerankerSimilarity::Cosine,
}
}
}
+19
View File
@@ -0,0 +1,19 @@
use anyhow::Result;
use candle_core::{DType, Device};
use crate::models::qwen3_reranker::model::Qwen3RerankerBackend;
pub struct Qwen3RerankerModel {
backend: Qwen3RerankerBackend,
}
impl Qwen3RerankerModel {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let backend = Qwen3RerankerBackend::load(path, device, dtype)?;
Ok(Self { backend })
}
pub fn rerank(&mut self, query: &str, documents: &[String]) -> Result<Vec<f32>> {
self.backend.rerank(query, documents)
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod config;
pub mod generate;
pub mod model;
+51
View File
@@ -0,0 +1,51 @@
use anyhow::{Result, anyhow};
use candle_core::{DType, Device};
use crate::models::{
common::retrieval::cosine_similarity,
qwen3_embedding::generate::Qwen3EmbeddingModel,
qwen3_reranker::config::{Qwen3RerankerConfig, Qwen3RerankerSimilarity},
};
pub struct Qwen3RerankerBackend {
config: Qwen3RerankerConfig,
embedding_backend: Qwen3EmbeddingModel,
}
impl Qwen3RerankerBackend {
pub fn load(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let embedding_backend = Qwen3EmbeddingModel::init(path, device, dtype)?;
Ok(Self {
config: Qwen3RerankerConfig::default(),
embedding_backend,
})
}
pub fn rerank(&mut self, query: &str, documents: &[String]) -> Result<Vec<f32>> {
if query.trim().is_empty() {
return Err(anyhow!("reranker query cannot be empty"));
}
if documents.is_empty() {
return Err(anyhow!("reranker documents cannot be empty"));
}
let mut batch = Vec::with_capacity(documents.len() + 1);
batch.push(query.to_string());
batch.extend(documents.iter().cloned());
let embeddings = self.embedding_backend.embed(&batch)?;
let query_embedding = embeddings
.first()
.ok_or_else(|| anyhow!("failed to produce query embedding"))?;
let mut scores = Vec::with_capacity(documents.len());
for doc_embedding in embeddings.iter().skip(1) {
scores.push(match self.config.similarity {
Qwen3RerankerSimilarity::Cosine => {
cosine_similarity(query_embedding, doc_embedding)?
}
});
}
Ok(scores)
}
}
+264
View File
@@ -0,0 +1,264 @@
use std::{
fs::File,
io::BufReader,
path::{Path, PathBuf},
};
use aha::models::{
common::{gguf::Gguf, retrieval::cosine_similarity},
qwen3_embedding::generate::Qwen3EmbeddingModel,
};
use anyhow::{Context, Result, anyhow};
use candle_core::{Device, quantized::gguf_file};
use ort::session::Session;
const QWEN3_EMBEDDING_SAFETENSORS_DIR: &str = r"D:\model_download\Qwen3-Embedding-0.6B";
const QWEN3_EMBEDDING_GGUF_DIR: &str = r"D:\model_download\Qwen3-Embedding-0.6B-GGUF";
const QWEN3_EMBEDDING_ONNX_DIR: &str = r"D:\model_download\Qwen3-Embedding-0.6B-ONNX";
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<PathBuf> {
require_existing_dir(dir)?;
let mut candidates = 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))
})
.collect::<Vec<_>>();
candidates.sort();
candidates
.into_iter()
.next()
.ok_or_else(|| anyhow!("no .{} file found in {}", extension, dir))
}
fn first_file_with_extension_recursive(dir: &str, extension: &str) -> Result<PathBuf> {
require_existing_dir(dir)?;
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(&current)? {
let entry = entry?;
let path = entry.path();
if path.is_dir() {
stack.push(path);
continue;
}
if path
.extension()
.is_some_and(|ext| ext.eq_ignore_ascii_case(extension))
{
matches.push(path);
}
}
}
matches.sort();
matches
.into_iter()
.next()
.ok_or_else(|| anyhow!("no .{} file found (recursive) in {}", extension, dir))
}
#[test]
fn qwen3_embedding_safetensors_can_load() -> Result<()> {
// Run this test only:
// cargo test --test test_qwen3_embedding_multi_format qwen3_embedding_safetensors_can_load -- --nocapture
require_existing_dir(QWEN3_EMBEDDING_SAFETENSORS_DIR)?;
let _model = Qwen3EmbeddingModel::init(QWEN3_EMBEDDING_SAFETENSORS_DIR, None, None)
.with_context(|| {
format!(
"failed to init safetensors model from {}",
QWEN3_EMBEDDING_SAFETENSORS_DIR
)
})?;
Ok(())
}
#[test]
fn qwen3_embedding_gguf_can_load() -> Result<()> {
// Run this test only:
// cargo test --test test_qwen3_embedding_multi_format qwen3_embedding_gguf_can_load -- --nocapture
let gguf_path = first_file_with_extension(QWEN3_EMBEDDING_GGUF_DIR, "gguf")?;
let file = File::open(&gguf_path)
.with_context(|| format!("failed to open gguf file: {}", gguf_path.display()))?;
let mut reader = BufReader::new(file);
let content = gguf_file::Content::read(&mut reader)
.with_context(|| format!("failed to parse gguf file: {}", gguf_path.display()))?;
if content.tensor_infos.is_empty() {
return Err(anyhow!(
"gguf tensor_infos is empty: {}",
gguf_path.display()
));
}
let gguf = Gguf::new(content, reader, Device::Cpu);
let tokenizer = gguf
.build_tokenizer(Some(false), Some(false), Some(false))
.context("failed to build tokenizer from gguf metadata")?;
let vocab_size = tokenizer.tokenizer.get_vocab_size(false);
if vocab_size == 0 {
return Err(anyhow!("gguf tokenizer vocab is empty"));
}
Ok(())
}
#[test]
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.
let metadata = std::fs::metadata(&onnx_path)
.with_context(|| format!("failed to read onnx metadata: {}", onnx_path.display()))?;
if metadata.len() == 0 {
return Err(anyhow!("onnx file is empty: {}", onnx_path.display()));
}
let _bytes = std::fs::read(&onnx_path)
.with_context(|| format!("failed to read onnx file: {}", onnx_path.display()))?;
Ok(())
}
#[test]
fn qwen3_embedding_onnxruntime_can_create_session() -> Result<()> {
// Run this test only:
// cargo test --test test_qwen3_embedding_multi_format qwen3_embedding_onnxruntime_can_create_session -- --nocapture
let onnx_path = first_file_with_extension_recursive(QWEN3_EMBEDDING_ONNX_DIR, "onnx")?;
// 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");
return Ok(());
}
let session = Session::builder()
.context("failed to create onnxruntime session builder")?
.commit_from_file(&onnx_path)
.with_context(|| {
format!(
"failed to create onnxruntime session from {}",
onnx_path.display()
)
})?;
if session.inputs().is_empty() {
return Err(anyhow!("onnxruntime session has no inputs"));
}
if session.outputs().is_empty() {
return Err(anyhow!("onnxruntime session has no outputs"));
}
Ok(())
}
#[test]
fn qwen3_embedding_real_texts_similarity() -> Result<()> {
// Run this test only:
// cargo test --test test_qwen3_embedding_multi_format qwen3_embedding_real_texts_similarity -- --exact --nocapture --test-threads=1
require_existing_dir(QWEN3_EMBEDDING_SAFETENSORS_DIR)?;
let query = "如何在 Rust 项目中做异步 HTTP 请求";
let documents = vec![
"Rust 中可以使用 reqwest + tokio 发起异步 HTTP 请求".to_string(),
"今天天气很好,适合出去散步和拍照".to_string(),
"在 Python 里可以用 requests 发送同步网络请求".to_string(),
"数据库索引优化可以显著提升查询性能".to_string(),
];
let mut model = Qwen3EmbeddingModel::init(QWEN3_EMBEDDING_SAFETENSORS_DIR, None, None)?;
let mut inputs = vec![query.to_string()];
inputs.extend(documents.clone());
let embeddings = model.embed(&inputs)?;
if embeddings.len() != inputs.len() {
return Err(anyhow!(
"embedding count mismatch: got {}, expect {}",
embeddings.len(),
inputs.len()
));
}
for (idx, emb) in embeddings.iter().enumerate() {
if emb.len() != 1024 {
return Err(anyhow!(
"embedding dim mismatch at index {}: got {}, expect 1024",
idx,
emb.len()
));
}
}
for (idx, text) in inputs.iter().enumerate() {
println!("text[{idx}]: {text}");
println!("embedding[{idx}] dim={}", embeddings[idx].len());
println!(
"embedding[{idx}]={}",
serde_json::to_string(&embeddings[idx])?
);
}
let query_embedding = &embeddings[0];
let mut similarities = Vec::with_capacity(documents.len());
let mut best_idx = 0usize;
let mut best_score = f32::NEG_INFINITY;
for (doc_idx, doc_emb) in embeddings.iter().enumerate().skip(1) {
let score = cosine_similarity(query_embedding, doc_emb)?;
similarities.push(score);
println!(
"similarity(query, doc_{}) = {:.6}, doc = {}",
doc_idx - 1,
score,
documents[doc_idx - 1]
);
if score > best_score {
best_score = score;
best_idx = doc_idx - 1;
}
}
println!("best_match_doc_index={}", best_idx);
println!("best_match_doc={}", documents[best_idx]);
println!("best_match_score={:.6}", best_score);
let result_json = serde_json::json!({
"query": query,
"documents": documents,
"embeddings": embeddings,
"similarities": similarities,
"best_match_doc_index": best_idx,
"best_match_doc": documents[best_idx],
"best_match_score": best_score
});
let output_path = Path::new("target").join("qwen3_embedding_similarity_output.json");
std::fs::write(&output_path, serde_json::to_string_pretty(&result_json)?)?;
println!("result_json_saved_to={}", output_path.display());
Ok(())
}