diff --git a/Cargo.lock b/Cargo.lock index 3c56c66..cb71c46 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -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" diff --git a/Cargo.toml b/Cargo.toml index f3b6db0..30d8cf1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -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"] } diff --git a/dev_rule.md b/dev_rule.md new file mode 100644 index 0000000..24e2778 --- /dev/null +++ b/dev_rule.md @@ -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//`: + +1. `src/models//` + 实现模型本体、配置、推理入口。 +2. `src/models/mod.rs` + 注册模块、枚举 `WhichModel`、模型元信息、模型工厂 `load_model`、`ModelInstance` 能力分派。 +3. `src/exec/.rs` + 接入 `aha run -m ` 的直调能力。 +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//` 的最低结构如下: + +```text +src/models// +├── 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// +├── 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// +├── 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// +├── 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// +├── 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// +├── 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 类型命名 + +- 配置类型:`Config` +- 生成配置类型:`GenerationConfig` +- 底层模型:`Model` 或 `Backend` +- 对外入口:`GenerateModel`、`EmbeddingModel`、`RerankerModel` +- 处理器:`Processor` + +### 5.3 `mod.rs` 命名要求 + +必须导出到以下位置: + +- `src/models/mod.rs` +- `src/exec/mod.rs` + +禁止: + +- 同一个模型目录里同时存在多个对外主入口但没有清晰职责说明。 + +--- + +## 6. 新增模型的强制清单 + +### 6.1 模型目录 + +必须完成: + +1. 创建 `src/models//` +2. 按模板补齐最少文件 +3. 保持职责边界清晰 + +### 6.2 模型注册 + +必须同时修改 `src/models/mod.rs`: + +1. `pub mod ;` +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/.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_.rs` +- 多格式专项:`tests/test__multi_format.rs` + +### 7.4 校验命令 + +最少校验: + +```bash +cargo fmt --check +cargo check +``` + +建议校验: + +```bash +cargo test --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// +├── mod.rs +├── config.rs +├── model.rs +└── generate.rs +``` + +### 10.2 原生多模态模型 + +```text +src/models// +├── mod.rs +├── config.rs +├── model.rs +├── generate.rs +└── processor.rs +``` + +### 10.3 多格式检索模型 + +```text +src/models// +├── mod.rs +├── config.rs +├── model.rs +├── generate.rs +├── backend_safetensors.rs # 可选 +├── backend_gguf.rs # 可选 +└── backend_onnx.rs # 可选 +``` + +最终原则只有一句: + +新增模型不是“把代码放进一个目录”,而是“把模型完整接入到 models factory、CLI、API、tests、docs 五条链路中”,并且目录职责必须与现有标准结构保持一致。 diff --git a/docs/supported-models.md b/docs/supported-models.md index 436ba73..457b3a7 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -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. diff --git a/docs/supported-models.zh-CN.md b/docs/supported-models.zh-CN.md index f24305d..5fecea7 100644 --- a/docs/supported-models.zh-CN.md +++ b/docs/supported-models.zh-CN.md @@ -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) 了解添加新模型集成的说明。 diff --git a/scripts/download_and_run.sh b/scripts/download_and_run.sh index 428638b..fb06643 100755 --- a/scripts/download_and_run.sh +++ b/scripts/download_and_run.sh @@ -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" ;; diff --git a/src/api/mod.rs b/src/api/mod.rs index 73074b9..69ccfe8 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -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) -> (Status, Stri } } +#[derive(Debug, Deserialize)] +pub(crate) struct EmbeddingRequest { + pub model: Option, + pub input: Value, +} + +#[derive(Debug, Serialize)] +struct EmbeddingData { + object: String, + index: usize, + embedding: Vec, +} + +#[derive(Debug, Serialize)] +struct EmbeddingResponse { + object: String, + data: Vec, + model: String, +} + +#[derive(Debug, Deserialize)] +pub(crate) struct RerankRequest { + pub model: Option, + pub query: String, + pub documents: Vec, + pub top_n: Option, +} + +#[derive(Debug, Serialize)] +struct RerankResult { + index: usize, + relevance_score: f32, + document: String, +} + +#[derive(Debug, Serialize)] +struct RerankResponse { + object: String, + model: String, + results: Vec, +} + +fn parse_embedding_input(input: &Value) -> anyhow::Result> { + 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 = "")] +pub(crate) async fn embeddings(req: Json) -> (Status, Json) { + 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::>(); + let response = EmbeddingResponse { + object: "list".to_string(), + data, + model: model_name, + }; + (Status::Ok, Json(serde_json::to_value(response).unwrap())) +} + +#[post("/rerank", data = "")] +pub(crate) async fn rerank(req: Json) -> (Status, Json) { + 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::>(); + 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)) { if let Some(model_ref) = MODEL.get() { @@ -307,10 +436,10 @@ pub(crate) async fn models() -> (Status, (ContentType, Json)) 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"); } } diff --git a/src/exec/fun_asr_nano.rs b/src/exec/fun_asr_nano.rs index 1af00fe..a69e870 100644 --- a/src/exec/fun_asr_nano.rs +++ b/src/exec/fun_asr_nano.rs @@ -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://") diff --git a/src/exec/glm_asr_nano.rs b/src/exec/glm_asr_nano.rs index 70c1d3a..8320700 100644 --- a/src/exec/glm_asr_nano.rs +++ b/src/exec/glm_asr_nano.rs @@ -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://") diff --git a/src/exec/mod.rs b/src/exec/mod.rs index 1e2c08e..fa28273 100644 --- a/src/exec/mod.rs +++ b/src/exec/mod.rs @@ -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; diff --git a/src/exec/qwen2_5vl.rs b/src/exec/qwen2_5vl.rs index ecaaa2e..1bd7138 100644 --- a/src/exec/qwen2_5vl.rs +++ b/src/exec/qwen2_5vl.rs @@ -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://") diff --git a/src/exec/qwen3_5.rs b/src/exec/qwen3_5.rs index c69e5a8..7c7dd16 100644 --- a/src/exec/qwen3_5.rs +++ b/src/exec/qwen3_5.rs @@ -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://") diff --git a/src/exec/qwen3_embedding.rs b/src/exec/qwen3_embedding.rs new file mode 100644 index 0000000..1be90d6 --- /dev/null +++ b/src/exec/qwen3_embedding.rs @@ -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(()) + } +} diff --git a/src/exec/qwen3_reranker.rs b/src/exec/qwen3_reranker.rs new file mode 100644 index 0000000..ddf1b3d --- /dev/null +++ b/src/exec/qwen3_reranker.rs @@ -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: " + )); + } + 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::>(); + 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> { + 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::>(); + Ok(docs) +} + +fn read_documents_file(path: &str) -> Result> { + 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::>(); + Ok(docs) +} diff --git a/src/main.rs b/src/main.rs index 99ae64e..b080503 100644 --- a/src/main.rs +++ b/src/main.rs @@ -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 = models + let model_infos: Vec = 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, + save_dir: Option, + download_retries: Option, + gguf_path: Option, + mmproj_path: Option, + allow_download: bool, +) -> anyhow::Result<(String, Option, Option)> { + 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 diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 8d871c2..e9b785b 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -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}, diff --git a/src/models/common/retrieval.rs b/src/models/common/retrieval.rs new file mode 100644 index 0000000..d896194 --- /dev/null +++ b/src/models/common/retrieval.rs @@ -0,0 +1,41 @@ +use anyhow::{Result, anyhow}; + +pub trait TextEmbeddingBackend { + fn embed_texts(&mut self, input: &[String]) -> Result>>; +} + +pub fn l2_normalize(v: &mut [f32]) { + let norm = v.iter().map(|x| x * x).sum::().sqrt(); + if norm > 0.0 { + for x in v.iter_mut() { + *x /= norm; + } + } +} + +pub fn mean_pool(embeddings: &[Vec]) -> Result> { + 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 { + if lhs.len() != rhs.len() { + return Err(anyhow!("embedding dimension mismatch")); + } + Ok(lhs.iter().zip(rhs.iter()).map(|(l, r)| l * r).sum::()) +} diff --git a/src/models/mod.rs b/src/models/mod.rs index dc3fe87..390cee2 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -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>), @@ -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>> { + 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> { + 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")); } diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs index 9543ad0..bbff76c 100644 --- a/src/models/qwen3/model.rs +++ b/src/models/qwen3/model.rs @@ -205,7 +205,11 @@ pub struct Qwen3Model { impl Qwen3Model { pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { - 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 { + 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 { 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 { Ok(self.embed_tokens.forward(input_ids)?) diff --git a/src/models/qwen3_embedding/config.rs b/src/models/qwen3_embedding/config.rs new file mode 100644 index 0000000..6898ff4 --- /dev/null +++ b/src/models/qwen3_embedding/config.rs @@ -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 { + 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, + }) + } +} diff --git a/src/models/qwen3_embedding/generate.rs b/src/models/qwen3_embedding/generate.rs new file mode 100644 index 0000000..238772a --- /dev/null +++ b/src/models/qwen3_embedding/generate.rs @@ -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) -> Result { + let backend = Qwen3EmbeddingBackend::load(path, device, dtype)?; + Ok(Self { backend }) + } + + pub fn embed(&mut self, input: &[String]) -> Result>> { + self.backend.embed_texts(input) + } +} + +impl TextEmbeddingBackend for Qwen3EmbeddingModel { + fn embed_texts(&mut self, input: &[String]) -> Result>> { + self.backend.embed_texts(input) + } +} diff --git a/src/models/qwen3_embedding/mod.rs b/src/models/qwen3_embedding/mod.rs new file mode 100644 index 0000000..7fca417 --- /dev/null +++ b/src/models/qwen3_embedding/mod.rs @@ -0,0 +1,3 @@ +pub mod config; +pub mod generate; +pub mod model; diff --git a/src/models/qwen3_embedding/model.rs b/src/models/qwen3_embedding/model.rs new file mode 100644 index 0000000..734f7bb --- /dev/null +++ b/src/models/qwen3_embedding/model.rs @@ -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) -> Result { + 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>> { + 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> { + 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::()?; + let mut pooled = match self.pooling { + Qwen3EmbeddingPoolingStrategy::Mean => mean_pool(&hidden_vec)?, + }; + if self.normalize { + l2_normalize(&mut pooled); + } + Ok(pooled) + } +} diff --git a/src/models/qwen3_reranker/config.rs b/src/models/qwen3_reranker/config.rs new file mode 100644 index 0000000..ed719bd --- /dev/null +++ b/src/models/qwen3_reranker/config.rs @@ -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, + } + } +} diff --git a/src/models/qwen3_reranker/generate.rs b/src/models/qwen3_reranker/generate.rs new file mode 100644 index 0000000..9d9aec3 --- /dev/null +++ b/src/models/qwen3_reranker/generate.rs @@ -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) -> Result { + let backend = Qwen3RerankerBackend::load(path, device, dtype)?; + Ok(Self { backend }) + } + + pub fn rerank(&mut self, query: &str, documents: &[String]) -> Result> { + self.backend.rerank(query, documents) + } +} diff --git a/src/models/qwen3_reranker/mod.rs b/src/models/qwen3_reranker/mod.rs new file mode 100644 index 0000000..7fca417 --- /dev/null +++ b/src/models/qwen3_reranker/mod.rs @@ -0,0 +1,3 @@ +pub mod config; +pub mod generate; +pub mod model; diff --git a/src/models/qwen3_reranker/model.rs b/src/models/qwen3_reranker/model.rs new file mode 100644 index 0000000..dcb5e88 --- /dev/null +++ b/src/models/qwen3_reranker/model.rs @@ -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) -> Result { + 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> { + 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) + } +} diff --git a/tests/test_qwen3_embedding_multi_format.rs b/tests/test_qwen3_embedding_multi_format.rs new file mode 100644 index 0000000..0f9c050 --- /dev/null +++ b/tests/test_qwen3_embedding_multi_format.rs @@ -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 { + 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::>(); + + 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 { + 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(()) +}