添加 rerank 和 embedding模型支持,并且添加 onnx 模型
This commit is contained in:
Generated
+60
@@ -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"
|
||||
|
||||
@@ -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
@@ -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 五条链路中”,并且目录职责必须与现有标准结构保持一致。
|
||||
@@ -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.
|
||||
|
||||
@@ -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) 了解添加新模型集成的说明。
|
||||
|
||||
@@ -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
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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://")
|
||||
|
||||
@@ -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://")
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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://")
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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},
|
||||
|
||||
@@ -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
@@ -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"));
|
||||
}
|
||||
|
||||
@@ -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)?)
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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(¤t)? {
|
||||
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(())
|
||||
}
|
||||
Reference in New Issue
Block a user