diff --git a/Cargo.lock b/Cargo.lock index 1d529ca..f7ea864 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -30,7 +30,7 @@ dependencies = [ [[package]] name = "aha" -version = "0.2.0" +version = "0.2.1" dependencies = [ "aha_openai_dive", "anyhow", diff --git a/Cargo.toml b/Cargo.toml index 4ae629b..d2fc8d2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,10 +1,10 @@ [package] name = "aha" -version = "0.2.0" +version = "0.2.1" edition = "2024" repository = "https://github.com/jhqxxx/aha" license = "Apache-2.0" -description = "aha model inference library, now supports Qwen2.5VL, MiniCPM4, VoxCPM, Qwen3VL, DeepSeek-OCR, Hunyuan-OCR, PaddleOCR-VL, VoxCPM1.5, RMBG2.0, GLM-ASR-Nano-2512, Fun-ASR-Nano-2512, Qwen3, Qwen3-ASR" +description = "aha model inference library, now supports Qwen2.5VL, MiniCPM4, VoxCPM, Qwen3VL, DeepSeek-OCR, Hunyuan-OCR, PaddleOCR-VL, VoxCPM1.5, RMBG2.0, GLM-ASR-Nano-2512, Fun-ASR-Nano-2512, Qwen3, Qwen3-ASR, Qwen3.5" [dependencies] candle-core = { version = "0.9.2" } diff --git a/README.md b/README.md index 17e6c57..3284aba 100644 --- a/README.md +++ b/README.md @@ -25,6 +25,9 @@ aha is a high-performance, cross-platform AI inference engine built with Rust and the Candle framework. It brings state-of-the-art AI models to your local machine—no API keys, no cloud dependencies, just pure, fast AI running directly on your hardware. ## Changelog +### v0.2.1 (2026-03-05) +- Added Qwen3.5 model + ### 2026-03-01 - update interpolate.rs diff --git a/README.zh-CN.md b/README.zh-CN.md index a3c72c2..9a6221a 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -25,6 +25,14 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理引擎。将最先进的 AI 模型带到您的本地机器——无需 API 密钥,无需云依赖,纯粹、快速的 AI,直接在您的硬件上运行。 ## 更新日志 +### v0.2.1 (2026-03-05) +- 新增Qwen3.5 模型 + +### 2026-03-01 +- 更新 interpolate.rs + +### 2026-02-24 +- 更新 candle 版本 0.9.2 ### v0.2.0 (2026-02-05) - 新增 Qwen3-ASR 语音识别模型 diff --git a/docs/changelog.md b/docs/changelog.md index 3b0267b..6044da0 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -5,6 +5,15 @@ All notable changes to aha will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [0.2.1] - (2026-03-05) +- Added Qwen3.5 model + +### 2026-03-01 +- update interpolate.rs + +### 2026-02-24 +- update candle version 0.9.2 + ## [0.2.0] - 2026-02-05 ### Added diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md index 24ba10d..a8d5500 100644 --- a/docs/changelog.zh-CN.md +++ b/docs/changelog.zh-CN.md @@ -5,6 +5,15 @@ 格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/), 本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。 +## [0.2.1] - (2026-03-05) +- 新增Qwen3.5 模型 + +### 2026-03-01 +- 更新 interpolate.rs + +### 2026-02-24 +- 更新 candle 版本 0.9.2 + ## [0.2.0] - 2026-02-05 ### 新增 diff --git a/docs/supported-models.md b/docs/supported-models.md index c39d597..00cc460 100644 --- a/docs/supported-models.md +++ b/docs/supported-models.md @@ -2,34 +2,27 @@ aha supports a growing collection of state-of-the-art AI models across multiple domains. -## Text Generation +## Language Model | Model | Parameters | Description | Use Case | License | |-------|-----------|-------------|----------|---------| -| **Qwen2.5-VL-3B** | 3B | Multimodal LLM | Chat, reasoning, vision | [Qwen Research License](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | -| **Qwen2.5-VL-7B** | 7B | Multimodal LLM | Chat, reasoning, vision | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **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) | ## Vision & Multimodal -| Model | Parameters | Description | Resolution | License | -|-------|-----------|-------------|------------|---------| -| **Qwen2.5-VL-3B** | 3B | Image understanding | Up to 1024x1024 | [Qwen Research License](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | -| **Qwen2.5-VL-7B** | 7B | Image understanding | Up to 1024x1024 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-2B** | 2B | Enhanced multimodal | Up to 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-4B** | 4B | Enhanced multimodal | Up to 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-8B** | 8B | Enhanced multimodal | Up to 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-32B** | 32B | Enhanced multimodal | Up to 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | - -## Speech Recognition (ASR) - -| Model | Parameters | Language | Real-time | Speed | License | -|-------|-----------|----------|-----------|-------|---------| -| **Fun-ASR-Nano-2512** | 2512M | Chinese/English | Yes | 16x realtime | Not Specified | -| **GLM-ASR-Nano-2512** | 2512M | Chinese/English | Yes | 32x realtime | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | -| **Qwen3-ASR-0.6B** | 0.6B | Chinese/English | Yes | Fast | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-ASR-1.7B** | 1.7B | Chinese/English | Yes | Fast | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| Model | Parameters | Description | License | +|-------|-----------|-------------|---------| +| **Qwen2.5-VL-3B** | 3B | Image understanding | [Qwen Research License](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | +| **Qwen2.5-VL-7B** | 7B | Image understanding | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-2B** | 2B | Enhanced multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-4B** | 4B | Enhanced multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-8B** | 8B | Enhanced multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-32B** | 32B | Enhanced multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3.5-0.8B** | 0.8B | Native Multimodal | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **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) | ## OCR @@ -39,6 +32,16 @@ aha supports a growing collection of state-of-the-art AI models across multiple | **Hunyuan-OCR** | Chinese | Deep learning | Complex layouts | [Tencent Hunyuan Community License](https://huggingface.co/tencent/HunyuanOCR/blob/main/LICENSE) | | **DeepSeek-OCR** | Multi | Scene text | Natural images | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | +## Speech Recognition (ASR) + +| Model | Parameters | Language | Real-time | Speed | License | +|-------|-----------|----------|-----------|-------|---------| +| **Fun-ASR-Nano-2512** | 2G | Chinese/English | Yes | Fast | Not Specified | +| **GLM-ASR-Nano-2512** | 4.5G | Chinese/English | Yes | Fast | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | +| **Qwen3-ASR-0.6B** | 0.6B | Chinese/English | Yes | Fast | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-ASR-1.7B** | 1.7B | Chinese/English | Yes | Fast | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | + + ## Audio Generation | Model | Parameters | Type | Description | License | @@ -77,7 +80,7 @@ Always verify license terms before deployment in production environments. ## Model Updates -Models are regularly updated. Check the [releases](https://github.com/jhqxxx/aha/releases) for the latest versions. +Models updated from time to time. ## Performance Benchmarks diff --git a/docs/supported-models.zh-CN.md b/docs/supported-models.zh-CN.md index f158db2..94bca43 100644 --- a/docs/supported-models.zh-CN.md +++ b/docs/supported-models.zh-CN.md @@ -6,30 +6,23 @@ aha 支持多个领域的最先进 AI 模型集合。 | 模型 | 参数量 | 描述 | 使用场景 | 开源协议 | |------|--------|------|----------|---------| -| **Qwen2.5-VL-3B** | 3B | 多模态大语言模型 | 对话、推理、视觉 | [Qwen 研究许可协议](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | -| **Qwen2.5-VL-7B** | 7B | 多模态大语言模型 | 对话、推理、视觉 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | | **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) | ## 视觉与多模态 -| 模型 | 参数量 | 描述 | 分辨率 | 开源协议 | +| 模型 | 参数量 | 描述 | 开源协议 | |------|--------|------|--------|---------| -| **Qwen2.5-VL-3B** | 3B | 图像理解 | 最高 1024x1024 | [Qwen 研究许可协议](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | -| **Qwen2.5-VL-7B** | 7B | 图像理解 | 最高 1024x1024 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-2B** | 2B | 增强多模态 | 最高 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-4B** | 4B | 增强多模态 | 最高 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-8B** | 8B | 增强多模态 | 最高 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-VL-32B** | 32B | 增强多模态 | 最高 1536x1536 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | - -## 语音识别 (ASR) - -| 模型 | 参数量 | 语言 | 实时 | 速度 | 开源协议 | -|------|--------|------|------|------|---------| -| **Fun-ASR-Nano-2512** | 2512M | 中/英 | 是 | 16x 实时 | 未标明 | -| **GLM-ASR-Nano-2512** | 2512M | 中/英 | 是 | 32x 实时 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | -| **Qwen3-ASR-0.6B** | 0.6B | 中/英 | 是 | 快速 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | -| **Qwen3-ASR-1.7B** | 1.7B | 中/英 | 是 | 快速 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen2.5-VL-3B** | 3B | 图像理解 | [Qwen 研究许可协议](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE) | +| **Qwen2.5-VL-7B** | 7B | 图像理解 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-2B** | 2B | 增强多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-4B** | 4B | 增强多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-8B** | 8B | 增强多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-VL-32B** | 32B | 增强多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3.5-0.8B** | 0.8B | 原生多模态 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **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) | ## OCR @@ -39,6 +32,16 @@ aha 支持多个领域的最先进 AI 模型集合。 | **Hunyuan-OCR** | 中文 | 深度学习 | 复杂布局 | [腾讯混元社区许可协议](https://huggingface.co/tencent/HunyuanOCR/blob/main/LICENSE) | | **DeepSeek-OCR** | 多语言 | 场景文字 | 自然图像 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | +## 语音识别 (ASR) + +| 模型 | 参数量 | 语言 | 实时 | 速度 | 开源协议 | +|------|--------|------|------|------|---------| +| **Fun-ASR-Nano-2512** | 2512M | 中/英 | 是 | 快速 | 未标明 | +| **GLM-ASR-Nano-2512** | 2512M | 中/英 | 是 | 快速 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | +| **Qwen3-ASR-0.6B** | 0.6B | 中/英 | 是 | 快速 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | +| **Qwen3-ASR-1.7B** | 1.7B | 中/英 | 是 | 快速 | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) | + + ## 语音生成 | 模型 | 参数量 | 类型 | 描述 | 开源协议 | @@ -77,7 +80,7 @@ aha 支持多个领域的最先进 AI 模型集合。 ## 模型更新 -模型定期更新。查看 [releases](https://github.com/jhqxxx/aha/releases) 获取最新版本。 +模型不定期更新。 ## 性能基准 diff --git a/src/api.rs b/src/api.rs index a4243b3..060324a 100644 --- a/src/api.rs +++ b/src/api.rs @@ -235,6 +235,10 @@ fn which_model_to_id(which_model: WhichModel) -> &'static str { 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::Qwen3ASR0_6B => "qwen3asr-0.6b", WhichModel::Qwen3ASR1_7B => "qwen3asr-1.7b", WhichModel::Qwen3vl2B => "qwen3vl-2b", @@ -262,6 +266,10 @@ fn which_model_to_owner(which_model: WhichModel) -> &'static str { | WhichModel::Qwen3vl4B | WhichModel::Qwen3vl8B | WhichModel::Qwen3vl32B => "Qwen", + WhichModel::Qwen3_5_0_8B + | WhichModel::Qwen3_5_2B + | WhichModel::Qwen3_5_4B + | WhichModel::Qwen3_5_9B => "Qwen", WhichModel::DeepSeekOCR => "deepseek-ai", WhichModel::HunyuanOCR => "Tencent-Hunyuan", WhichModel::PaddleOCRVL => "PaddlePaddle", diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs index 3a4c3d2..df269d0 100644 --- a/src/chat_template/mod.rs +++ b/src/chat_template/mod.rs @@ -131,4 +131,25 @@ impl<'a> ChatTemplate<'a> { .map_err(|e| anyhow!(format!("render template error {}", e)))?; Ok(message_str) } + + pub fn apply_chat_temp_think( + &self, + messages: &ChatCompletionParameters, + enable_thinking: Option, + ) -> Result { + let context = context! { + messages => &messages.messages, + tools => &messages.tools.as_ref(), + add_generation_prompt => true, + enable_thinking => enable_thinking, + }; + let template = self + .env + .get_template("chat") + .map_err(|e| anyhow!(format!("render template error {}", e)))?; + let message_str = template + .render(context) + .map_err(|e| anyhow!(format!("render template error {}", e)))?; + Ok(message_str) + } } diff --git a/src/exec/mod.rs b/src/exec/mod.rs index 5eb4708..aac7164 100644 --- a/src/exec/mod.rs +++ b/src/exec/mod.rs @@ -11,6 +11,7 @@ pub mod minicpm4; pub mod paddleocr_vl; pub mod qwen2_5vl; pub mod qwen3; +pub mod qwen3_5; pub mod qwen3_asr; pub mod qwen3vl; pub mod rmbg2_0; diff --git a/src/exec/qwen3_5.rs b/src/exec/qwen3_5.rs new file mode 100644 index 0000000..fd0cd86 --- /dev/null +++ b/src/exec/qwen3_5.rs @@ -0,0 +1,103 @@ +//! Qwen3.5 exec implementation for CLI `run` subcommand + +use std::time::Instant; + +use anyhow::Result; + +use crate::exec::ExecModel; +use crate::models::GenerateModel; +use crate::models::qwen3_5::generate::Qwen3_5GenerateModel; +use crate::utils::get_file_path; + +pub struct Qwen3_5Exec; + +impl ExecModel for Qwen3_5Exec { + fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> { + let input_text = &input[0]; + let target_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 = 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 input_url = if url.starts_with("http://") + || url.starts_with("https://") + || url.starts_with("file://") + { + url.clone() + } else { + format!("file://{}", url) + }; + let message = if input_url.ends_with("mp4") { + format!( + r#"{{ + "model": "qwen3.5", + "messages": [ + {{ + "role": "user", + "content": [ + {{ + "type": "video", + "video_url": + {{ + "url": "{}" + }} + }}, + {{ + "type": "text", + "text": "{}" + }} + ] + }} + ] + }}"#, + input_url, target_text + ) + } else { + format!( + r#"{{ + "model": "qwen2.5", + "messages": [ + {{ + "role": "user", + "content": [ + {{ + "type": "image", + "image_url": {{ + "url": "{}" + }} + }}, + {{ + "type": "text", + "text": "{}" + }} + ] + }} + ] + }}"#, + input_url, target_text + ) + }; + let mes = serde_json::from_str(&message)?; + + let i_start = Instant::now(); + let result = model.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + + println!("Result: {:?}", result); + + if let Some(out) = output { + std::fs::write(out, format!("{:?}", result))?; + println!("Output saved to: {}", out); + } + + Ok(()) + } +} diff --git a/src/exec/qwen3vl.rs b/src/exec/qwen3vl.rs index a336d24..493a170 100644 --- a/src/exec/qwen3vl.rs +++ b/src/exec/qwen3vl.rs @@ -61,7 +61,7 @@ impl ExecModel for Qwen3vlExec { } else { format!( r#"{{ - "model": "qwen2.5vl", + "model": "qwen3vl", "messages": [ {{ "role": "user", diff --git a/src/lib.rs b/src/lib.rs index 0d6e9ef..7fdf7eb 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -5,3 +5,5 @@ pub mod position_embed; pub mod process; pub mod tokenizer; pub mod utils; + +pub use aha_openai_dive::v1::resources::chat::{ChatCompletionParameters, ChatCompletionResponse}; diff --git a/src/main.rs b/src/main.rs index 92ff3f4..95d8179 100644 --- a/src/main.rs +++ b/src/main.rs @@ -201,6 +201,10 @@ fn run_list(args: ListArgs) -> anyhow::Result<()> { 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, @@ -390,6 +394,22 @@ fn run_run(args: RunArgs) -> anyhow::Result<()> { use aha::exec::qwen3::Qwen3Exec; Qwen3Exec::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)?; + } + WhichModel::Qwen3_5_2B => { + use aha::exec::qwen3_5::Qwen3_5Exec; + Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?; + } + WhichModel::Qwen3_5_4B => { + use aha::exec::qwen3_5::Qwen3_5Exec; + Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?; + } + WhichModel::Qwen3_5_9B => { + use aha::exec::qwen3_5::Qwen3_5Exec; + Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?; + } WhichModel::Qwen3ASR0_6B => { use aha::exec::qwen3_asr::Qwen3ASRExec; Qwen3ASRExec::run(&input, output.as_deref(), &weight_path)?; diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 080c861..05abc32 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -1,4 +1,4 @@ -use anyhow::{Result, anyhow}; +use anyhow::Result; use candle_core::{D, IndexOp, Tensor}; use candle_nn::{ Activation, BatchNorm, BatchNormConfig, Conv1d, Conv1dConfig, Conv2d, Conv2dConfig, @@ -6,7 +6,6 @@ use candle_nn::{ ModuleT, RmsNorm, VarBuilder, batch_norm, conv1d, conv1d_no_bias, conv2d, conv2d_no_bias, embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm, }; -use rayon::iter::{IndexedParallelIterator, IntoParallelRefIterator, ParallelIterator}; use crate::{ position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer}, @@ -971,49 +970,6 @@ impl LlamaForCausalLM { } } -pub fn conv1d_group_parallel(xs: &Tensor, conv1d: &Conv1d) -> Result { - let groups = conv1d.config().groups; - let xs = if groups == 1 { - xs.conv1d_with_algo( - conv1d.weight(), - conv1d.config().padding, - conv1d.config().stride, - conv1d.config().dilation, - groups, - conv1d.config().cudnn_fwd_algo, - )? - } else { - let blocks = xs.chunk(groups, 1)?; - let kernel = conv1d.weight().chunk(groups, 0)?; - let blocks = blocks - // .iter() - .par_iter() - .zip(&kernel) - .map(|(block, kernel)| { - block - .conv1d_with_algo( - kernel, - conv1d.config().padding, - conv1d.config().stride, - conv1d.config().dilation, - 1, - conv1d.config().cudnn_fwd_algo, - ) - .map_err(|e| anyhow!(format!("tensor conv1d_with_algo error:{}", e))) - }) - .collect::>>()?; - Tensor::cat(&blocks, 1)? - }; - match conv1d.bias() { - None => Ok(xs), - Some(bias) => { - let b = bias.dims1()?; - let bias = bias.reshape((1, b, 1))?; - Ok(xs.broadcast_add(&bias)?) - } - } -} - pub struct GLU { dim: usize, } @@ -1194,3 +1150,45 @@ pub fn mish(xs: &Tensor) -> Result { let xs = xs.mul(&tanh)?; Ok(xs) } + +pub fn softplus(xs: &Tensor) -> Result { + // ln(1 + exp(x)) + Ok((xs.exp()? + 1.0)?.log()?) +} + +pub fn softplus_stable(xs: &Tensor) -> Result { + // max(x, 0) + ln(1 + exp(-abs(x))) + let zero = Tensor::zeros_like(xs)?; + let x_max_0 = xs.maximum(&zero)?; + Ok((xs.abs()?.neg()?.exp()? + 1.0)?.log()?.add(&x_max_0)?) +} + +// refer to https://github.com/huggingface/candle/issues/3389 +pub fn conv1d_depthwise(input: &Tensor, weight: &Tensor, bias: Option<&Tensor>) -> Result { + // group = dim, stride= 1 + // input: (bs, dim, len) + // weight: (dim, 1, k) -> (dim, k) + // input already padding + let len_in = input.dim(2)?; + let weight = weight.squeeze(1)?; + let kernel_size = weight.dim(1)?; + // len_out = (len_in - k + 2p) / s + 1, p = 0, s = 1 + let len_out = len_in - kernel_size + 1; + let mut out = input + .narrow(2, 0, len_out)? + .broadcast_mul(&weight.narrow(1, 0, 1)?.unsqueeze(0)?)?; + for k in 1..kernel_size { + out = (out + + input + .narrow(2, k, len_out)? + .broadcast_mul(&weight.narrow(1, k, 1)?.unsqueeze(0)?)?)?; + } + match bias { + None => Ok(out), + Some(bias) => { + let b = bias.dims1()?; + let bias = bias.reshape((1, b, 1))?; + Ok(out.broadcast_add(&bias)?) + } + } +} diff --git a/src/models/fun_asr_nano/model.rs b/src/models/fun_asr_nano/model.rs index 779cc38..fa17903 100644 --- a/src/models/fun_asr_nano/model.rs +++ b/src/models/fun_asr_nano/model.rs @@ -5,7 +5,8 @@ use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::soft use crate::{ models::{ common::{ - NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm, + NaiveAttention, TwoLinearMLP, conv1d_depthwise, eager_attention_forward, get_conv1d, + get_layer_norm, }, fun_asr_nano::config::FunASRNanoConfig, qwen3::{config::Qwen3Config, model::Qwen3Model}, @@ -85,7 +86,8 @@ impl MultiHeadedAttentionSANM { }; let xs = inputs.transpose(1, 2)?; let xs = xs.pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)?; - let xs = self.fsmn_block.forward(&xs)?; + // let xs = self.fsmn_block.forward(&xs)?; + let xs = conv1d_depthwise(&xs, self.fsmn_block.weight(), self.fsmn_block.bias())?; let xs = xs.transpose(1, 2)?; let mut xs = xs.add(&inputs)?; if let Some(mask) = mask { diff --git a/src/models/mask_gct/model.rs b/src/models/mask_gct/model.rs index 51c66c5..e7cac03 100644 --- a/src/models/mask_gct/model.rs +++ b/src/models/mask_gct/model.rs @@ -6,11 +6,10 @@ use candle_nn::{ use crate::{ models::{ - common::{WNConv1d, get_conv1d, get_layer_norm}, + common::{WNConv1d, conv1d_depthwise, get_conv1d, get_layer_norm}, mask_gct::config::SemanticCodec, }, - utils::interpolate::interpolate_nearest_1d, - utils::tensor_utils::l2_normalize, + utils::{interpolate::interpolate_nearest_1d, tensor_utils::l2_normalize}, }; pub struct ConvNeXtBlock { @@ -43,7 +42,9 @@ impl ConvNeXtBlock { } pub fn forward(&self, xs: &Tensor) -> Result { let residual = xs.clone(); - let xs = self.dwconv.forward(xs)?; + // let xs = self.dwconv.forward(xs)?; + let xs = xs.pad_with_zeros(D::Minus1, 3, 3)?; + let xs = conv1d_depthwise(&xs, self.dwconv.weight(), self.dwconv.bias())?; let xs = xs.transpose(1, 2)?; let xs = self.norm.forward(&xs)?; let xs = self.pwconv1.forward(&xs)?.gelu()?; diff --git a/src/models/mod.rs b/src/models/mod.rs index b1f00f9..741d629 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -11,6 +11,7 @@ pub mod minicpm4; pub mod paddleocr_vl; pub mod qwen2_5vl; pub mod qwen3; +pub mod qwen3_5; pub mod qwen3_asr; pub mod qwen3vl; pub mod rmbg2_0; @@ -29,9 +30,9 @@ use crate::models::{ glm_asr_nano::generate::GlmAsrNanoGenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, - qwen3::generate::Qwen3GenerateModel, qwen3_asr::generate::Qwen3AsrGenerateModel, - qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model, - voxcpm::generate::VoxCPMGenerate, + qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel, + qwen3_asr::generate::Qwen3AsrGenerateModel, qwen3vl::generate::Qwen3VLGenerateModel, + rmbg2_0::generate::RMBG2_0Model, voxcpm::generate::VoxCPMGenerate, }; #[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)] @@ -44,6 +45,14 @@ pub enum WhichModel { Qwen2_5vl7B, #[value(name = "qwen3-0.6b", hide = true)] Qwen3_0_6B, + #[value(name = "qwen3.5-0.8b", hide = true)] + Qwen3_5_0_8B, + #[value(name = "qwen3.5-2b", hide = true)] + Qwen3_5_2B, + #[value(name = "qwen3.5-4b", hide = true)] + Qwen3_5_4B, + #[value(name = "qwen3.5-9b", hide = true)] + Qwen3_5_9B, #[value(name = "qwen3asr-0.6b", hide = true)] Qwen3ASR0_6B, #[value(name = "qwen3asr-1.7b", hide = true)] @@ -82,6 +91,10 @@ 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::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::Qwen3ASR0_6B => "Qwen/Qwen3-ASR-0.6B", WhichModel::Qwen3ASR1_7B => "Qwen/Qwen3-ASR-1.7B", WhichModel::Qwen3vl2B => "Qwen/Qwen3-VL-2B-Instruct", @@ -103,14 +116,17 @@ impl WhichModel { pub fn model_type(self) -> &'static str { match self { // LLM models - WhichModel::MiniCPM4_0_5B - | WhichModel::Qwen2_5vl3B + WhichModel::MiniCPM4_0_5B | WhichModel::Qwen3_0_6B => "llm", + WhichModel::Qwen2_5vl3B | WhichModel::Qwen2_5vl7B - | WhichModel::Qwen3_0_6B | WhichModel::Qwen3vl2B | WhichModel::Qwen3vl4B | WhichModel::Qwen3vl8B - | WhichModel::Qwen3vl32B => "llm", + | WhichModel::Qwen3vl32B + | WhichModel::Qwen3_5_0_8B + | WhichModel::Qwen3_5_2B + | WhichModel::Qwen3_5_4B + | WhichModel::Qwen3_5_9B => "vlm", // OCR models WhichModel::DeepSeekOCR | WhichModel::HunyuanOCR | WhichModel::PaddleOCRVL => "ocr", // ASR models @@ -143,6 +159,7 @@ pub enum ModelInstance<'a> { MiniCPM4(MiniCPMGenerateModel<'a>), Qwen2_5VL(Qwen2_5VLGenerateModel<'a>), Qwen3(Qwen3GenerateModel<'a>), + Qwen3_5(Qwen3_5GenerateModel<'a>), Qwen3ASR(Qwen3AsrGenerateModel<'a>), Qwen3VL(Qwen3VLGenerateModel<'a>), DeepSeekOCR(DeepseekOCRGenerateModel), @@ -160,6 +177,7 @@ 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::Qwen3_5(model) => model.generate(mes), ModelInstance::Qwen3ASR(model) => model.generate(mes), ModelInstance::Qwen3VL(model) => model.generate(mes), ModelInstance::DeepSeekOCR(model) => model.generate(mes), @@ -187,6 +205,7 @@ 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::Qwen3_5(model) => model.generate_stream(mes), ModelInstance::Qwen3VL(model) => model.generate_stream(mes), ModelInstance::Qwen3ASR(model) => model.generate_stream(mes), ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes), @@ -218,6 +237,22 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result { + let model = Qwen3_5GenerateModel::init(path, None, None)?; + ModelInstance::Qwen3_5(model) + } + WhichModel::Qwen3_5_2B => { + let model = Qwen3_5GenerateModel::init(path, None, None)?; + ModelInstance::Qwen3_5(model) + } + WhichModel::Qwen3_5_4B => { + let model = Qwen3_5GenerateModel::init(path, None, None)?; + ModelInstance::Qwen3_5(model) + } + WhichModel::Qwen3_5_9B => { + let model = Qwen3_5GenerateModel::init(path, None, None)?; + ModelInstance::Qwen3_5(model) + } WhichModel::Qwen3ASR0_6B => { let model = Qwen3AsrGenerateModel::init(path, None, None)?; ModelInstance::Qwen3ASR(model) diff --git a/src/models/qwen3_5/config.rs b/src/models/qwen3_5/config.rs new file mode 100644 index 0000000..29d4a6b --- /dev/null +++ b/src/models/qwen3_5/config.rs @@ -0,0 +1,57 @@ +use candle_nn::Activation; + +use crate::models::qwen3vl::config::Qwen3VLVisionConfig; + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct RopeParameters { + pub mrope_interleaved: bool, + pub mrope_section: Vec, + pub rope_type: String, + pub rope_theta: f32, + pub partial_rotary_factor: f32, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3_5TextConfig { + pub attention_bias: bool, + pub attention_dropout: f32, + pub attn_output_gate: bool, + pub dtype: String, + pub eos_token_id: u32, + pub full_attention_interval: usize, + pub head_dim: usize, + pub hidden_act: Activation, + pub hidden_size: usize, + pub initializer_range: f32, + pub intermediate_size: usize, + pub layer_types: Vec, + pub linear_conv_kernel_dim: usize, + pub linear_key_head_dim: usize, + pub linear_num_key_heads: usize, + pub linear_num_value_heads: usize, + pub linear_value_head_dim: usize, + pub max_position_embeddings: usize, + pub mlp_only_layers: Vec, + pub mtp_num_hidden_layers: usize, + pub mtp_use_dedicated_embeddings: bool, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + pub rms_norm_eps: f64, + pub tie_word_embeddings: bool, + pub use_cache: bool, + pub vocab_size: usize, + pub mamba_ssm_dtype: String, + pub rope_parameters: RopeParameters, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3_5Config { + pub image_token_id: u32, + pub text_config: Qwen3_5TextConfig, + pub tie_word_embeddings: bool, + pub video_token_id: u32, + pub vision_config: Qwen3VLVisionConfig, + pub vision_end_token_id: u32, + pub vision_start_token_id: u32, +} diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs new file mode 100644 index 0000000..b4a989f --- /dev/null +++ b/src/models/qwen3_5/generate.rs @@ -0,0 +1,231 @@ +use aha_openai_dive::v1::resources::chat::{ + ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, +}; +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, Tensor}; +use candle_nn::VarBuilder; +use rocket::async_stream::stream; +use rocket::futures::Stream; + +use crate::{ + chat_template::ChatTemplate, + models::{ + GenerateModel, + qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model}, + qwen3vl::processor::Qwen3VLProcessor, + }, + tokenizer::TokenizerModel, + utils::{ + build_completion_chunk_response, build_completion_response, extract_metadata_value, find_type_files, get_device, get_dtype, get_logit_processor + }, +}; + +pub struct Qwen3_5GenerateModel<'a> { + chat_template: ChatTemplate<'a>, + tokenizer: TokenizerModel, + pre_processor: Qwen3VLProcessor, + qwen3_5: Qwen3_5Model, + device: Device, + eos_token_id: u32, + model_name: String, +} + +impl<'a> Qwen3_5GenerateModel<'a> { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let chat_template = ChatTemplate::init(path)?; + let tokenizer = TokenizerModel::init(path)?; + let config_path = path.to_string() + "/config.json"; + let cfg: Qwen3_5Config = serde_json::from_slice(&std::fs::read(config_path)?)?; + let device = get_device(device); + let cfg_dtype = cfg.text_config.dtype.as_str(); + let dtype = get_dtype(dtype, cfg_dtype); + let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?; + let model_list = find_type_files(path, "safetensors")?; + let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; + let eos_token_id = cfg.text_config.eos_token_id; + let qwen3_5 = Qwen3_5Model::new(vb, cfg)?; + + Ok(Self { + chat_template, + tokenizer, + pre_processor, + qwen3_5, + device, + eos_token_id, + model_name: "qwen3.5".to_string(), + }) + } +} + +impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking"); + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); + let mes_render = self.chat_template.apply_chat_temp_think(&mes, enable_thinking)?; + let input = self.pre_processor.process_info(&mes, &mes_render)?; + let mut input_ids = self + .tokenizer + .text_encode(input.replace_text.clone(), &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let mut pixel_values = input.pixel_values.as_ref(); + let image_grid_thw = input.image_grid_thw.as_ref(); + let mut pixel_values_video = input.pixel_values_video.as_ref(); + let video_grid_thw = input.video_grid_thw.as_ref(); + let mut generate = Vec::new(); + let sample_len = mes.max_tokens.unwrap_or(1024); + for _ in 0..sample_len { + let logits = self.qwen3_5.forward( + &input_ids, + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + generate.push(next_token); + if next_token == self.eos_token_id { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + pixel_values = None; + pixel_values_video = None; + } + let num_token = generate.len() as u32; + let res = self.tokenizer.token_decode(generate)?; + self.qwen3_5.clear_cache(); + let response = build_completion_response(res, &self.model_name, Some(num_token)); + Ok(response) + } + + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result< + Box< + dyn Stream> + + Send + + Unpin + + '_, + >, + > { + let seed = match mes.seed { + None => 34562u64, + Some(s) => s as u64, + }; + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); + let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking"); + let mes_render = self.chat_template.apply_chat_temp_think(&mes, enable_thinking)?; + let input = self.pre_processor.process_info(&mes, &mes_render)?; + let mut input_ids = self + .tokenizer + .text_encode(input.replace_text.clone(), &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let sample_len = mes.max_tokens.unwrap_or(1024); + let stream = stream! { + let mut error_tokens = Vec::new(); + let mut pixel_values = input.pixel_values.as_ref(); + let image_grid_thw = input.image_grid_thw.as_ref(); + let mut pixel_values_video = input.pixel_values_video.as_ref(); + let video_grid_thw = input.video_grid_thw.as_ref(); + let mut tool_call_id = None; + let mut tool_call_content = String::new(); + for _ in 0..sample_len { + let logits = self.qwen3_5.forward( + &input_ids, + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + let mut decode_ids = Vec::new(); + if !error_tokens.is_empty() { + decode_ids.extend_from_slice(&error_tokens); + } + decode_ids.push(next_token); + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?; + if decoded_token.contains("�") { + error_tokens.push(next_token); + if error_tokens.len() > 3 { + error_tokens.clear(); + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + pixel_values = None; + pixel_values_video = None; + continue; + } + error_tokens.clear(); + // 处理特殊标记和工具调用 + match decoded_token.as_str() { + "" => { + // 开始工具调用 + tool_call_id = Some(uuid::Uuid::new_v4().to_string()); + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + pixel_values = None; + pixel_values_video = None; + continue; + } + "" => { + // 结束工具调用 + let chunk = build_completion_chunk_response( + decoded_token, + &self.model_name, + tool_call_id.clone(), + Some(tool_call_content.clone()) + ); + tool_call_id = None; + tool_call_content = String::new(); + yield Ok(chunk); + } + _ => { + if tool_call_id.is_some() { + // 在工具调用过程中,收集工具调用内容 + tool_call_content.push_str(&decoded_token); + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + pixel_values = None; + pixel_values_video = None; + continue; + } else { + // 正常文本输出 + let chunk = build_completion_chunk_response( + decoded_token, + &self.model_name, + None, + None + ); + yield Ok(chunk); + } + } + } + if next_token == self.eos_token_id { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + pixel_values = None; + pixel_values_video = None; + } + self.qwen3_5.clear_cache(); + }; + Ok(Box::new(Box::pin(stream))) + } +} diff --git a/src/models/qwen3_5/mod.rs b/src/models/qwen3_5/mod.rs new file mode 100644 index 0000000..7fca417 --- /dev/null +++ b/src/models/qwen3_5/mod.rs @@ -0,0 +1,3 @@ +pub mod config; +pub mod generate; +pub mod model; diff --git a/src/models/qwen3_5/model.rs b/src/models/qwen3_5/model.rs new file mode 100644 index 0000000..0255261 --- /dev/null +++ b/src/models/qwen3_5/model.rs @@ -0,0 +1,1141 @@ +use anyhow::{Result, anyhow}; +use candle_core::{D, IndexOp, Tensor}; +use candle_nn::{ + Conv1d, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias, + ops::sigmoid, rms_norm, +}; + +use crate::{ + models::{ + common::{GateUpDownMLP, conv1d_depthwise, eager_attention_forward, get_conv1d, softplus}, + qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig}, + qwen3vl::model::Qwen3VLVisionModel, + }, + position_embed::rope::{Qwen3VLTextRotaryEmbedding, glm_asr_apply_rotary_pos_emb}, + utils::tensor_utils::{ + get_equal_mask, get_vision_next_indices, l2_normalize, masked_scatter_dim0, nonzero_index, + prepare_causal_attention_mask, repeat_interleave, split_tensor, zero_index, + }, +}; + +pub struct Qwen3_5RMSNorm { + eps: f64, + weight: Tensor, +} + +impl Qwen3_5RMSNorm { + pub fn new(vb: VarBuilder, dim: usize, eps: f64) -> Result { + let weight = vb.get(dim, "weight")?; + let weight = weight.to_dtype(candle_core::DType::F32)?.affine(1.0, 1.0)?; + Ok(Self { eps, weight }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let x = xs.to_dtype(candle_core::DType::F32)?; + let norm_ = x + .powf(2.0)? + .mean_keepdim(D::Minus1)? + .affine(1.0, self.eps)? + .sqrt()?; + let norm = x.broadcast_div(&norm_)?; + let norm = norm.broadcast_mul(&self.weight)?.to_dtype(xs.dtype())?; + Ok(norm) + } +} + +pub struct Qwen3_5RMSNormGated { + norm: RmsNorm, +} + +impl Qwen3_5RMSNormGated { + pub fn new(vb: VarBuilder, hidden_size: usize, eps: f64) -> Result { + let norm = rms_norm(hidden_size, eps, vb)?; + Ok(Self { norm }) + } + + pub fn forward(&self, xs: &Tensor, gate: Option<&Tensor>) -> Result { + let mut xs = self.norm.forward(xs)?; + if let Some(gate) = gate { + xs = xs.broadcast_mul(&gate.silu()?)?; + } + Ok(xs) + } +} + +macro_rules! transmute_tensors { + ($($tensor:expr),*) => { + ($( + $tensor.transpose(1, 2)?.contiguous()?.to_dtype(candle_core::DType::F32)?, + )*) + }; +} + +macro_rules! right_pad_zero_tensor { + ($dim:expr, $pad_size:expr, $($tensor:expr),+) => { + ($( + $tensor.pad_with_zeros($dim, 0, $pad_size)?.contiguous()?, + )+) + }; +} + +macro_rules! reshape_chunk_tensor { + ($chunk_size:expr, $($tensor:expr),*) => { + ($( + { + let (bs, head, _, dim) = $tensor.dims4()?; + $tensor.reshape((bs, head, (), $chunk_size, dim))?.contiguous()? + }, + )*) + }; +} + +pub struct Qwen3_5GatedDeltaNet { + // hidden_size: usize, + num_v_heads: usize, + num_k_heads: usize, + head_k_dim: usize, + head_v_dim: usize, + key_dim: usize, + value_dim: usize, + conv_kernel_size: usize, + // layer_idx: usize, + // activation: Activation, + // layer_norm_epsilon: f64, + // QKV 投影 + // conv_dim: usize, + conv1d: Conv1d, + dt_bias: Tensor, + a_log: Tensor, + norm: Qwen3_5RMSNormGated, + out_proj: Linear, + + // Z, B, A 投影 + in_proj_qkv: candle_nn::Linear, + in_proj_z: candle_nn::Linear, + in_proj_b: candle_nn::Linear, + in_proj_a: candle_nn::Linear, + conv_state_cache: Option, + recurrent_state_cache: Option, +} + +impl Qwen3_5GatedDeltaNet { + pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result { + let hidden_size = config.hidden_size; + let num_v_heads = config.linear_num_value_heads; + let num_k_heads = config.linear_num_key_heads; + let head_k_dim = config.linear_key_head_dim; + let head_v_dim = config.linear_value_head_dim; + let key_dim = head_k_dim * num_k_heads; + let value_dim = head_v_dim * num_v_heads; + let conv_kernel_size = config.linear_conv_kernel_dim; + // let activation = config.hidden_act; + let layer_norm_epsilon = config.rms_norm_eps; + let conv_dim = key_dim * 2 + value_dim; + let conv1d = get_conv1d( + vb.pp("conv1d"), + conv_dim, + conv_dim, + conv_kernel_size, + conv_kernel_size - 1, + 1, + 1, + conv_dim, + false, + )?; + let dt_bias = vb.get(num_v_heads, "dt_bias")?; + let a_log = vb.get(num_v_heads, "A_log")?; + let norm = Qwen3_5RMSNormGated::new(vb.pp("norm"), head_v_dim, layer_norm_epsilon)?; + + let out_proj = linear_no_bias(value_dim, hidden_size, vb.pp("out_proj"))?; + + let in_proj_qkv = linear_no_bias(hidden_size, conv_dim, vb.pp("in_proj_qkv"))?; + let in_proj_z = linear_no_bias(hidden_size, value_dim, vb.pp("in_proj_z"))?; + let in_proj_b = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_b"))?; + let in_proj_a = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_a"))?; + Ok(Self { + // hidden_size, + num_v_heads, + num_k_heads, + head_k_dim, + head_v_dim, + key_dim, + value_dim, + conv_kernel_size, + // activation, + // layer_norm_epsilon, + // conv_dim, + conv1d, + dt_bias, + a_log, + norm, + out_proj, + in_proj_qkv, + in_proj_z, + in_proj_b, + in_proj_a, + conv_state_cache: None, + recurrent_state_cache: None, + }) + } + + fn torch_causal_conv1d_update(&mut self, xs: &Tensor) -> Result { + let conv_state = self.conv_state_cache.as_ref().unwrap(); + let seq_len = xs.dim(2)?; + let state_len = conv_state.dim(D::Minus1)?; + let conv_state_new = Tensor::cat(&[conv_state, xs], D::Minus1)?; + let conv_update = conv_state_new.narrow(D::Minus1, seq_len, state_len)?; + self.conv_state_cache = Some(conv_update); + // too slow + // let out = conv_state_new.conv1d(self.conv1d.weight(), 0, 1, 1, dim)?; + let out = conv1d_depthwise(&conv_state_new, self.conv1d.weight(), self.conv1d.bias())?; + let start = out.dim(D::Minus1)? - seq_len; + let out = out.narrow(D::Minus1, start, seq_len)?.silu()?; + Ok(out) + } + + fn torch_chunk_gated_delta_rule( + &mut self, + query: &Tensor, + key: &Tensor, + value: &Tensor, + g: &Tensor, + beta: &Tensor, + use_qk_l2norm_in_kernel: bool, + chunk_size: usize, + ) -> Result { + let (query, key) = if use_qk_l2norm_in_kernel { + (l2_normalize(query, 3)?, l2_normalize(key, 3)?) + } else { + (query.clone(), key.clone()) + }; + let initial_dtype = query.dtype(); + let (query, key, value, beta, g) = transmute_tensors!(query, key, value, beta, g); + let (batch_size, num_heads, sequence_length, k_head_dim) = key.dims4()?; + let v_head_dim = value.dim(D::Minus1)?; + let pad_size = (chunk_size - sequence_length % chunk_size) % chunk_size; + let (query, key, value, beta, g) = + right_pad_zero_tensor!(2, pad_size, query, key, value, beta, g); + let total_sequence_length = sequence_length + pad_size; + let scale = 1.0 / (query.dim(D::Minus1)? as f64).sqrt(); + let query = query.affine(scale, 0.0)?; + let v_beta = value.broadcast_mul(&beta.unsqueeze(D::Minus1)?.contiguous()?)?; + let k_beta = key.broadcast_mul(&beta.unsqueeze(D::Minus1)?.contiguous()?)?; + let (query, key, k_beta, v_beta) = + reshape_chunk_tensor!(chunk_size, query, key, k_beta, v_beta); + let g = g.reshape((g.dim(0)?, g.dim(1)?, (), chunk_size))?; + let g = g.cumsum(D::Minus1)?; + let decay_mask = g + .unsqueeze(D::Minus1)? + .broadcast_sub(&g.unsqueeze(D::Minus2)?)? + .exp()? + .to_dtype(candle_core::DType::F32)?; + let tril_mask = Tensor::tril2(chunk_size, candle_core::DType::U32, query.device())? + .broadcast_as(decay_mask.shape())?; + let on_false = decay_mask.zeros_like()?; + let decay_mask = tril_mask.where_cond(&decay_mask, &on_false)?.contiguous()?; + // when rank = 5, matmul error, bs=1,squeeze(0) + let mut attn = k_beta + .squeeze(0)? + .contiguous()? + .matmul( + &key.squeeze(0)? + .transpose(D::Minus1, D::Minus2)? + .contiguous()?, + )? + .unsqueeze(0)? + .mul(&decay_mask)? + .affine(-1.0, 0.0)?; + // 包含对角线的上三角矩阵 + let mask = Tensor::triu2(chunk_size, candle_core::DType::U32, query.device())? + .broadcast_as(decay_mask.shape())?; + // attn的对角线为0,为1的取0, 为0的取attn + attn = mask.where_cond(&on_false, &attn)?; + let (d0, d1, d2, _, _) = attn.dims5()?; + for i in 1..chunk_size { + let row = attn.i((.., .., .., i, ..i))?.contiguous()?; + let sub = attn.i((.., .., .., ..i, ..i))?.contiguous()?; + let attn_i = row + .unsqueeze(D::Minus1)? + .broadcast_mul(&sub)? + .sum(D::Minus2)? + .add(&row)? + .unsqueeze(D::Minus2)?; + attn = attn.slice_assign(&[(0..d0), (0..d1), (0..d2), (i..i + 1), (0..i)], &attn_i)?; + } + let attn = attn + .broadcast_add(&Tensor::eye(chunk_size, attn.dtype(), attn.device())?)? + .contiguous()?; + // when rank = 5, matmul error, bs=1,squeeze(0) + let value = attn.squeeze(0)?.matmul(&v_beta.squeeze(0)?)?.unsqueeze(0)?; + let k_cumdecay = attn + .squeeze(0)? + .matmul( + &k_beta + .broadcast_mul(&g.exp()?.unsqueeze(D::Minus1)?)? + .squeeze(0)?, + )? + .unsqueeze(0)?; + let mut last_recurrent_state = if let Some(recurrent) = self.recurrent_state_cache.as_ref() + { + recurrent.clone() + } else { + Tensor::zeros( + (batch_size, num_heads, k_head_dim, v_head_dim), + candle_core::DType::F32, + value.device(), + )? + }; + + let mut core_attn_out = value.zeros_like()?; + let tril_mask = Tensor::tril2(chunk_size, candle_core::DType::U32, query.device())? + .broadcast_as((batch_size, num_heads, chunk_size, chunk_size))?; + let on_false = tril_mask.zeros_like()?.to_dtype(candle_core::DType::F32)?; + let last_dim = core_attn_out.dim(D::Minus1)?; + for i in 0..total_sequence_length / chunk_size { + let q_i = query.i((.., .., i))?.contiguous()?; + let k_i = key.i((.., .., i))?.contiguous()?; + let v_i = value.i((.., .., i))?.contiguous()?; + let g_i = g.i((.., .., i))?.contiguous()?; + let attn = q_i + .matmul(&k_i.transpose(D::Minus1, D::Minus2)?.contiguous()?)? + .mul(&decay_mask.i((.., .., i))?)?; + let attn = tril_mask.where_cond(&attn, &on_false)?.contiguous()?; + let v_prime = k_cumdecay.i((.., .., i))?.matmul(&last_recurrent_state)?; + let v_new = v_i.sub(&v_prime)?; + let attn_inter = q_i + .broadcast_mul(&g_i.unsqueeze(D::Minus1)?.exp()?)? + .matmul(&last_recurrent_state)?; + let out_i = attn_inter.add(&attn.matmul(&v_new)?)?.unsqueeze(2)?; + core_attn_out = core_attn_out.slice_assign( + &[ + (0..batch_size), + (0..num_heads), + (i..i + 1), + (0..chunk_size), + (0..last_dim), + ], + &out_i, + )?; + let g_i_last_dim = g_i.dim(D::Minus1)?; + last_recurrent_state = last_recurrent_state + .broadcast_mul( + &g_i.narrow(D::Minus1, g_i_last_dim - 1, 1)? + .unsqueeze(D::Minus1)? + .exp()?, + )? + .add( + &k_i.broadcast_mul( + &g_i.narrow(D::Minus1, g_i_last_dim - 1, 1)? + .broadcast_sub(&g_i)? + .exp()? + .unsqueeze(D::Minus1)?, + )? + .transpose(D::Minus1, D::Minus2)? + .squeeze(0)? + .matmul(&v_new.squeeze(0)?)? + .unsqueeze(0)?, + )?; + } + self.recurrent_state_cache = Some(last_recurrent_state); + core_attn_out = + core_attn_out.reshape((batch_size, num_heads, (), core_attn_out.dim(D::Minus1)?))?; + core_attn_out = core_attn_out.narrow(2, 0, sequence_length)?; + core_attn_out = core_attn_out + .transpose(1, 2)? + .contiguous()? + .to_dtype(initial_dtype)?; + + Ok(core_attn_out) + } + + fn torch_recurrent_gated_delta_rule( + &mut self, + query: &Tensor, + key: &Tensor, + value: &Tensor, + g: &Tensor, + beta: &Tensor, + use_qk_l2norm_in_kernel: bool, + ) -> Result { + let (query, key) = if use_qk_l2norm_in_kernel { + (l2_normalize(query, 3)?, l2_normalize(key, 3)?) + } else { + (query.clone(), key.clone()) + }; + let initial_dtype = query.dtype(); + let (query, key, value, beta, g) = transmute_tensors!(query, key, value, beta, g); + let (batch_size, num_heads, sequence_length, k_head_dim) = key.dims4()?; + let v_head_dim = value.dim(D::Minus1)?; + let scale = 1.0 / (query.dim(D::Minus1)? as f64).sqrt(); + let query = query.affine(scale, 0.0)?; + let mut last_recurrent_state = if let Some(recurrent) = self.recurrent_state_cache.as_ref() + { + recurrent.clone() + } else { + Tensor::zeros( + (batch_size, num_heads, k_head_dim, v_head_dim), + candle_core::DType::F32, + value.device(), + )? + }; + + let mut core_attn_out = Tensor::zeros( + (batch_size, num_heads, sequence_length, v_head_dim), + candle_core::DType::F32, + value.device(), + )?; + for i in 0..sequence_length { + let q_i = query.i((.., .., i))?; + let k_i = key.i((.., .., i))?; + let v_i = value.i((.., .., i))?; + let g_i = g + .i((.., .., i))? + .exp()? + .unsqueeze(D::Minus1)? + .unsqueeze(D::Minus1)?; + let beta_i = beta.i((.., .., i))?.unsqueeze(D::Minus1)?; + last_recurrent_state = last_recurrent_state.broadcast_mul(&g_i)?; + let kv_mem = last_recurrent_state + .broadcast_mul(&k_i.unsqueeze(D::Minus1)?)? + .sum(D::Minus2)?; + let delta = v_i.broadcast_sub(&kv_mem)?.broadcast_mul(&beta_i)?; + last_recurrent_state = last_recurrent_state.broadcast_add( + &k_i.unsqueeze(D::Minus1)? + .broadcast_mul(&delta.unsqueeze(D::Minus2)?)?, + )?; + let out_i = last_recurrent_state + .broadcast_mul(&q_i.unsqueeze(D::Minus1)?)? + .sum_keepdim(D::Minus2)?; + core_attn_out = core_attn_out.slice_assign( + &[(0..batch_size), (0..num_heads), (i..i + 1), (0..v_head_dim)], + &out_i, + )?; + } + self.recurrent_state_cache = Some(last_recurrent_state); + core_attn_out = core_attn_out + .transpose(1, 2)? + .contiguous()? + .to_dtype(initial_dtype)?; + + Ok(core_attn_out) + } + + pub fn forward(&mut self, xs: &Tensor, attention_mask: Option<&Tensor>) -> Result { + let xs = if let Some(mask) = attention_mask { + xs.broadcast_mul(&mask.unsqueeze(D::Minus1)?)? + } else { + xs.clone() + }; + let (bs, seq_len, _) = xs.dims3()?; + let mut mixed_qkv = self.in_proj_qkv.forward(&xs)?.transpose(1, 2)?; + let z = self + .in_proj_z + .forward(&xs)? + .reshape((bs, seq_len, (), self.head_v_dim))?; + let b = self.in_proj_b.forward(&xs)?; + let a = self.in_proj_a.forward(&xs)?; + let use_precomputed_states = + self.conv_state_cache.is_some() && self.recurrent_state_cache.is_some() && seq_len == 1; + if use_precomputed_states { + mixed_qkv = self.torch_causal_conv1d_update(&mixed_qkv)?; + } else { + let pad = self.conv_kernel_size as isize - mixed_qkv.dim(D::Minus1)? as isize; + let conv_state = if pad >= 0 { + mixed_qkv.pad_with_zeros(D::Minus1, pad as usize, 0)? + } else { + mixed_qkv.narrow(D::Minus1, pad.unsigned_abs(), self.conv_kernel_size)? + }; + self.conv_state_cache = Some(conv_state); + mixed_qkv = mixed_qkv.pad_with_zeros( + D::Minus1, + self.conv_kernel_size - 1, + self.conv_kernel_size - 1, + )?; + mixed_qkv = conv1d_depthwise(&mixed_qkv, self.conv1d.weight(), self.conv1d.bias())?; + mixed_qkv = mixed_qkv.narrow(D::Minus1, 0, seq_len)?.silu()?; + // too slowly + // mixed_qkv = self + // .conv1d + // .forward(&mixed_qkv)? + // .narrow(D::Minus1, 0, seq_len)? + // .silu()?; + } + let mixed_qkv = mixed_qkv.transpose(1, 2)?; + let qkv_split = split_tensor( + &mixed_qkv, + &[self.key_dim, self.key_dim, self.value_dim], + D::Minus1, + )?; + + let mut query = qkv_split[0].reshape((bs, seq_len, (), self.head_k_dim))?; + let mut key = qkv_split[1].reshape((bs, seq_len, (), self.head_k_dim))?; + let value = qkv_split[2].reshape((bs, seq_len, (), self.head_v_dim))?; + let beta = sigmoid(&b)?; + let a_plus_bias = softplus( + &a.to_dtype(candle_core::DType::F32)? + .broadcast_add(&self.dt_bias.to_dtype(candle_core::DType::F32)?)?, + )?; + let g = (-1.0 * self.a_log.to_dtype(candle_core::DType::F32)?.exp()?)? + .broadcast_mul(&a_plus_bias)?; + if self.num_v_heads / self.num_k_heads > 1 { + query = repeat_interleave(&query, self.num_v_heads / self.num_k_heads, 2)?; + key = repeat_interleave(&key, self.num_v_heads / self.num_k_heads, 2)?; + } + let core_attn_out = if !use_precomputed_states { + self.torch_chunk_gated_delta_rule(&query, &key, &value, &g, &beta, true, 64)? + } else { + self.torch_recurrent_gated_delta_rule(&query, &key, &value, &g, &beta, true)? + }; + let core_attn_out = core_attn_out.reshape(((), self.head_v_dim))?; + let z = z.reshape(((), self.head_v_dim))?; + let core_attn_out = self.norm.forward(&core_attn_out, Some(&z))?; + let core_attn_out = core_attn_out.reshape((bs, seq_len, ()))?; + let output = self.out_proj.forward(&core_attn_out)?; + + Ok(output) + } + + pub fn clear_cache(&mut self) { + self.conv_state_cache = None; + self.recurrent_state_cache = None; + } +} + +pub struct Qwen3_5Attention { + q_proj: Linear, + k_proj: Linear, + v_proj: Linear, + o_proj: Linear, + q_norm: Qwen3_5RMSNorm, + k_norm: Qwen3_5RMSNorm, + num_attention_heads: usize, + num_key_value_heads: usize, + num_kv_groups: usize, + head_dim: usize, + scaling: f64, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl Qwen3_5Attention { + pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result { + let hidden_size = config.hidden_size; + let num_attention_heads = config.num_attention_heads; + let head_dim = config.head_dim; + let num_key_value_heads = config.num_key_value_heads; + let num_kv_groups = num_attention_heads / num_key_value_heads; + let scaling = 1f64 / f64::sqrt(head_dim as f64); + let q_proj = linear_b( + hidden_size, + num_attention_heads * head_dim * 2, + config.attention_bias, + vb.pp("q_proj"), + )?; + let k_proj = linear_b( + hidden_size, + num_key_value_heads * head_dim, + config.attention_bias, + vb.pp("k_proj"), + )?; + let v_proj = linear_b( + hidden_size, + num_key_value_heads * head_dim, + config.attention_bias, + vb.pp("v_proj"), + )?; + let o_proj = linear_b( + num_attention_heads * head_dim, + hidden_size, + config.attention_bias, + vb.pp("o_proj"), + )?; + let q_norm = Qwen3_5RMSNorm::new(vb.pp("q_norm"), head_dim, config.rms_norm_eps)?; + let k_norm = Qwen3_5RMSNorm::new(vb.pp("k_norm"), head_dim, config.rms_norm_eps)?; + Ok(Self { + q_proj, + k_proj, + v_proj, + o_proj, + q_norm, + k_norm, + num_attention_heads, + num_key_value_heads, + num_kv_groups, + head_dim, + scaling, + kv_cache: None, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let (b_sz, q_len, _) = xs.dims3()?; + let query_chunk = self + .q_proj + .forward(xs)? + .reshape((b_sz, q_len, self.num_attention_heads, self.head_dim * 2))? + .chunk(2, D::Minus1)?; + let query_states = + query_chunk[0].reshape((b_sz, q_len, self.num_attention_heads, self.head_dim))?; + let gate = query_chunk[1].reshape((b_sz, q_len, ()))?; + + let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?; + let key_states = self.k_proj.forward(xs)?.reshape(( + b_sz, + q_len, + self.num_key_value_heads, + self.head_dim, + ))?; + let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?; + let value_states = self.v_proj.forward(xs)?; + let value_states = value_states + .reshape((b_sz, q_len, self.num_key_value_heads, self.head_dim))? + .transpose(1, 2)?; + let (query_states, key_states) = + glm_asr_apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?; + let (key_states, value_states) = match &self.kv_cache { + None => (key_states, value_states), + Some((prev_k, prev_v)) => { + let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; + let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; + (key_states, value_states) + } + }; + self.kv_cache = Some((key_states.clone(), value_states.clone())); + let attn_output = eager_attention_forward( + &query_states, + &key_states, + &value_states, + Some(self.num_kv_groups), + attention_mask, + self.scaling, + )?; + let attn_output = attn_output + .reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))? + .contiguous()?; + let attn_output = attn_output.mul(&sigmoid(&gate)?)?; + let attn_output = attn_output.apply(&self.o_proj)?; + Ok(attn_output) + } + + pub fn clear_kv_cache(&mut self) { + self.kv_cache = None + } +} + +pub struct Qwen3_5DecoderLayer { + // hidden_size: usize, + layer_type: String, + linear_attn: Option, + self_attn: Option, + mlp: GateUpDownMLP, + input_layernorm: Qwen3_5RMSNorm, + post_attention_layernorm: Qwen3_5RMSNorm, +} + +impl Qwen3_5DecoderLayer { + pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig, layer_idx: usize) -> Result { + let hidden_size = config.hidden_size; + let layer_type = config.layer_types[layer_idx].clone(); + let (linear_attn, self_attn) = if layer_type.eq("linear_attention") { + let linear_attn = Qwen3_5GatedDeltaNet::new(vb.pp("linear_attn"), config)?; + (Some(linear_attn), None) + } else { + let self_attn = Qwen3_5Attention::new(vb.pp("self_attn"), config)?; + (None, Some(self_attn)) + }; + let mlp = GateUpDownMLP::new( + vb.pp("mlp"), + hidden_size, + config.intermediate_size, + config.hidden_act, + false, + None, + None, + None, + )?; + let input_layernorm = + Qwen3_5RMSNorm::new(vb.pp("input_layernorm"), hidden_size, config.rms_norm_eps)?; + let post_attention_layernorm = Qwen3_5RMSNorm::new( + vb.pp("post_attention_layernorm"), + hidden_size, + config.rms_norm_eps, + )?; + Ok(Self { + // hidden_size, + layer_type, + linear_attn, + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: Option<&Tensor>, + sin: Option<&Tensor>, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let mut xs = self.input_layernorm.forward(xs)?; + if self.layer_type.eq("linear_attention") + && let Some(linear_attn) = self.linear_attn.as_mut() + { + xs = linear_attn.forward(&xs, attention_mask)?; + } else if let Some(self_attn) = self.self_attn.as_mut() + && let Some(cos) = cos + && let Some(sin) = sin + { + xs = self_attn.forward(&xs, cos, sin, attention_mask)?; + } + let residual = xs.add(&residual)?; + xs = self.post_attention_layernorm.forward(&residual)?; + xs = self.mlp.forward(&xs)?; + xs = xs.add(&residual)?; + Ok(xs) + } + + pub fn clear_cache(&mut self) { + if let Some(linear_attn) = self.linear_attn.as_mut() { + linear_attn.clear_cache(); + } + if let Some(self_attn) = self.self_attn.as_mut() { + self_attn.clear_kv_cache(); + } + } +} + +pub struct Qwen3_5TextModel { + embed_tokens: Embedding, + layers: Vec, + norm: Qwen3_5RMSNorm, + rotary_emb: Qwen3VLTextRotaryEmbedding, + mrope_section: Vec, +} + +impl Qwen3_5TextModel { + pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result { + let embed_tokens = embedding(config.vocab_size, config.hidden_size, vb.pp("embed_tokens"))?; + let mut layers = vec![]; + let vb_layers = vb.pp("layers"); + for i in 0..config.num_hidden_layers { + let layer = Qwen3_5DecoderLayer::new(vb_layers.pp(i), config, i)?; + layers.push(layer); + } + let norm = Qwen3_5RMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?; + let rope_dim = + (config.head_dim as f32 * config.rope_parameters.partial_rotary_factor) as usize; + let rotary_emb = + Qwen3VLTextRotaryEmbedding::new(rope_dim, config.rope_parameters.rope_theta); + Ok(Self { + embed_tokens, + layers, + norm, + rotary_emb, + mrope_section: config.rope_parameters.mrope_section.clone(), + }) + } + + pub fn forward(&mut self, inputs_embeds: &Tensor, position_ids: &Tensor) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + + let (cos, sin) = self.rotary_emb.forward( + position_ids, + inputs_embeds.dtype(), + self.mrope_section.clone(), + )?; + let mut xs = inputs_embeds.clone(); + let attention_mask: Option = { + if seq_len <= 1 { + None + } else { + Some(prepare_causal_attention_mask( + b_size, + seq_len, + 0, + inputs_embeds.device(), + )?) + } + }; + for layer in self.layers.iter_mut() { + let layer_mask = + if layer.layer_type.ne("linear_attention") || (seq_len != 1 && b_size != 1) { + attention_mask.clone() + } else { + None + }; + xs = layer.forward(&xs, Some(&cos), Some(&sin), layer_mask.as_ref())?; + } + xs = self.norm.forward(&xs)?; + Ok(xs) + } + + pub fn clear_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_cache(); + } + } +} + +pub struct Qwen3_5Model { + config: Qwen3_5Config, + visual: Qwen3VLVisionModel, + language_model: Qwen3_5TextModel, + lm_head: Linear, + rope_deltas: Option, +} + +impl Qwen3_5Model { + pub fn new(vb: VarBuilder, config: Qwen3_5Config) -> Result { + let vb_m = vb.pp("model"); + let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?; + let language_model = Qwen3_5TextModel::new(vb_m.pp("language_model"), &config.text_config)?; + let lm_head = if config.tie_word_embeddings { + Linear::new(language_model.embed_tokens.embeddings().clone(), None) + } else { + linear_no_bias( + config.text_config.hidden_size, + config.text_config.vocab_size, + vb.pp("lm_head"), + )? + }; + Ok(Self { + config, + visual, + language_model, + lm_head, + rope_deltas: None, + }) + } + + fn get_rope_index( + &self, + input_ids: &Tensor, + image_grid_thw: Option<&Tensor>, + video_grid_thw: Option<&Tensor>, + mask: Option<&Tensor>, + ) -> Result<(Tensor, Tensor)> { + let video_grid_thw = match video_grid_thw { + Some(thw) => { + let grid_t = thw.i((.., 0))?.to_vec1::()?; + let mut v_thw_vec = Vec::new(); + for (index, t) in grid_t.iter().enumerate() { + let mut thw_i = thw.i(index)?.to_vec1::()?; + // [12, 30, 50] + // [1, 30, 50]*t + thw_i[0] = 1; + v_thw_vec.push( + Tensor::new(thw_i, thw.device())? + .repeat(*t as usize)? + .reshape((*t as usize, ()))?, + ); + } + Some(Tensor::cat(&v_thw_vec, 0)?) + } + None => None, + }; + + let spatial_merge_size = self.config.vision_config.spatial_merge_size; + let image_token_id = self.config.image_token_id; + let video_token_id = self.config.video_token_id; + let vision_start_token_id = self.config.vision_start_token_id; + let mut mrope_position_deltas = vec![]; + if image_grid_thw.is_some() || video_grid_thw.is_some() { + let total_input_ids = input_ids.clone(); + let mask_ = mask + .cloned() + .unwrap_or(Tensor::ones_like(&total_input_ids)?) + .to_device(input_ids.device())?; + let mut position_ids = Tensor::ones( + (3, input_ids.dim(0)?, input_ids.dim(1)?), + input_ids.dtype(), + input_ids.device(), + )?; + let mut image_index = 0; + let mut video_index = 0; + + for i in 0..total_input_ids.dim(0)? { + let mut input_ids_i = total_input_ids.i(i)?; + let mask_i = mask_.i(i)?; + // 推理时, attention_mask如果是全1向量,取非0索引的操作没必要 + if mask_i.sum_all()?.to_scalar::()? != mask_i.dim(0)? as u32 { + let nonzero_idx = nonzero_index(&mask_i)?; + input_ids_i = input_ids_i.gather(&nonzero_idx, 0)?; + } + let mut text_start = 0; + let mut text_end = 0; + let mut thw = vec![]; + let mut llm_pos_ids_list: Vec = Vec::new(); + // vision start的下一个索引 + let vision_indices = get_vision_next_indices(&input_ids_i, vision_start_token_id); + + match vision_indices { + Ok(indeices) => { + let vision_tokens = input_ids_i.gather(&indeices, 0)?.to_vec1::()?; + let vision_indices_vec = indeices.to_vec1::()?; + for (j, &token) in vision_tokens.iter().enumerate() { + if token == image_token_id { + thw = image_grid_thw.unwrap().i(image_index)?.to_vec1::()?; + image_index += 1; + text_end = vision_indices_vec[j]; + } + if token == video_token_id { + thw = video_grid_thw + .as_ref() + .unwrap() + .i(video_index)? + .to_vec1::()?; + text_end = vision_indices_vec[j]; + video_index += 1; + } + let llm_grid_t = thw[0]; + let llm_grid_h = thw[1] / spatial_merge_size as u32; + let llm_grid_w = thw[2] / spatial_merge_size as u32; + let text_len = text_end - text_start; + let start_idx = if !llm_pos_ids_list.is_empty() { + llm_pos_ids_list[llm_pos_ids_list.len() - 1] + .max_all()? + .to_scalar::()? + + 1 + } else { + 0 + }; + let pos_ids = Tensor::arange( + start_idx, + start_idx + text_len, + input_ids_i.device(), + )? + .unsqueeze(0)? + .broadcast_as((3usize, text_len as usize))?; + llm_pos_ids_list.push(pos_ids); + + let t_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_t, + input_ids_i.device(), + )? + .unsqueeze(D::Minus1)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize * llm_grid_w as usize, + ))? + .flatten_all()?; + let h_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_h, + input_ids_i.device(), + )? + .unsqueeze(0)? + .unsqueeze(D::Minus1)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize, + llm_grid_w as usize, + ))? + .flatten_all()?; + let w_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_w, + input_ids_i.device(), + )? + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize, + llm_grid_w as usize, + ))? + .flatten_all()?; + + let thw_index = Tensor::stack(&[t_index, h_index, w_index], 0)?; + llm_pos_ids_list.push(thw_index); + text_start = text_end + llm_grid_t * llm_grid_h * llm_grid_w; + } + } + Err(e) => { + println!("get vision_indices err: {e}"); + } + }; + if text_start < input_ids_i.dim(0)? as u32 { + let start_idx = if !llm_pos_ids_list.is_empty() { + llm_pos_ids_list[llm_pos_ids_list.len() - 1] + .max_all()? + .to_scalar::()? + + 1 + } else { + 0 + }; + let text_len = input_ids_i.dim(0)? as u32 - text_start; + let pos_ids = + Tensor::arange(start_idx, start_idx + text_len, input_ids_i.device())? + .unsqueeze(0)? + .broadcast_as((3usize, text_len as usize))?; + llm_pos_ids_list.push(pos_ids); + } + let llm_position = Tensor::cat(&llm_pos_ids_list, 1)?.reshape((3, 1, ()))?; + position_ids = position_ids + .slice_assign(&[(0..3), (i..i + 1), (0..input_ids.dim(1)?)], &llm_position)?; + let position_deltas = llm_position.max_all()?.to_scalar::()? as i64 + 1 + - input_ids_i.dim(0)? as i64; + mrope_position_deltas.push(position_deltas); + } + let mut mrope_position_deltas = Tensor::new(mrope_position_deltas, input_ids.device())?; + if mrope_position_deltas.rank() == 1 { + mrope_position_deltas = mrope_position_deltas.unsqueeze(0)?; + } + Ok((position_ids.contiguous()?, mrope_position_deltas)) + } else if let Some(mask) = mask { + let mut position_ids = mask + .to_dtype(candle_core::DType::F64)? + .cumsum(D::Minus1)? + .to_dtype(candle_core::DType::U32)? + .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; + for i in 0..position_ids.dim(0)? { + let mut position_ids_i = position_ids.i(i)?; + let mask_i = mask.i(i)?; + // 如果有pad, 将填充位置置为1 + // 当bs>1, 可能存在不同序列长度,需要添加pad使seq_len长度一致 + if mask_i.sum_all()?.to_scalar::()? != mask_i.dim(0)? as u32 { + let zero_indices = zero_index(&mask_i)?; + let replace_1 = Tensor::ones( + zero_indices.dim(0)?, + candle_core::DType::U32, + input_ids.device(), + )?; + position_ids_i = position_ids_i + .scatter(&zero_indices, &replace_1, 0)? + .unsqueeze(0)?; + position_ids = position_ids + .slice_assign(&[(i..i + 1), (0..position_ids.dim(1)?)], &position_ids_i)?; + } + } + position_ids = position_ids + .unsqueeze(0)? + .broadcast_as((3, input_ids.dim(0)?, input_ids.dim(1)?))? + .contiguous()?; + let mut mrope_position_deltas = position_ids + .max(0)? + .max(D::Minus1)? + .broadcast_sub(&Tensor::new( + vec![mask.dim(D::Minus1)? as u32 - 1], + input_ids.device(), + )?)? + .contiguous()?; + if mrope_position_deltas.rank() == 1 { + mrope_position_deltas = mrope_position_deltas.unsqueeze(0)?; + } + Ok((position_ids, mrope_position_deltas)) + } else { + let position_ids = + Tensor::arange(0_u32, input_ids.dim(D::Minus1)? as u32, input_ids.device())? + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as((3, input_ids.dim(0)?, input_ids.dim(D::Minus1)?))? + .contiguous()?; + let mrope_position_deltas = Tensor::zeros( + (input_ids.dim(0)?, 1), + input_ids.dtype(), + input_ids.device(), + )?; + Ok((position_ids, mrope_position_deltas)) + } + } + + fn compute_3d_position_ids( + &mut self, + input_ids: &Tensor, + inputs_embeds: &Tensor, + image_grid_thw: Option<&Tensor>, + video_grid_thw: Option<&Tensor>, + seqlen_offset: usize, + ) -> Result { + let position_ids = if self.rope_deltas.is_none() { + let (position_ids, rope_deltas) = + self.get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)?; + self.rope_deltas = Some(rope_deltas); + position_ids + } else { + let (bs, seq_len, _) = inputs_embeds.dims3()?; + Tensor::arange( + seqlen_offset as i64, + (seqlen_offset + seq_len) as i64, + input_ids.device(), + )? + .unsqueeze(0)? + .broadcast_as((bs, seq_len))? + .broadcast_add(self.rope_deltas.as_ref().unwrap())? + .unsqueeze(0)? + .broadcast_as((3, bs, seq_len))? + .contiguous()? + .to_dtype(candle_core::DType::U32)? + }; + Ok(position_ids) + } + + pub fn forward( + &mut self, + input_ids: &Tensor, + pixel_values: Option<&Tensor>, + image_grid_thw: Option<&Tensor>, + pixel_values_video: Option<&Tensor>, + video_grid_thw: Option<&Tensor>, + // cache_position: Option<&Tensor>, + seqlen_offset: usize, + ) -> Result { + let mut inputs_embeds = self.language_model.embed_tokens.forward(input_ids)?; + if let Some(pixel_values) = pixel_values + && let Some(image_grid_thw) = image_grid_thw + { + let (image_embeds, _) = self.visual.forward(pixel_values, image_grid_thw)?; + let vision_mask = get_equal_mask(input_ids, self.config.image_token_id)?; + let n_image_tokens = vision_mask.sum_all()?.to_scalar::()?; + if n_image_tokens as usize != image_embeds.dim(0)? { + return Err(anyhow!(format!( + "n_image_token num: {} not equal to image_embed len: {}", + n_image_tokens, + image_embeds.dim(0)? + ))); + } + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &image_embeds, &vision_mask)?; + } + if let Some(pixel_values_video) = pixel_values_video + && let Some(video_grid_thw) = video_grid_thw + { + let (video_embeds, _) = self.visual.forward(pixel_values_video, video_grid_thw)?; + let vision_mask = get_equal_mask(input_ids, self.config.video_token_id)?; + let n_video_tokens = vision_mask.sum_all()?.to_scalar::()?; + if n_video_tokens as usize != video_embeds.dim(0)? { + return Err(anyhow!(format!( + "n_video_tokens num: {} not equal to video_embeds len: {}", + n_video_tokens, + video_embeds.dim(0)? + ))); + } + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &video_embeds, &vision_mask)?; + } + + let position_ids = self.compute_3d_position_ids( + input_ids, + &inputs_embeds, + image_grid_thw, + video_grid_thw, + seqlen_offset, + )?; + let outputs = self.language_model.forward(&inputs_embeds, &position_ids)?; + let seq_len = outputs.dim(1)?; + let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&hidden_state)?; + Ok(logits) + } + + pub fn clear_cache(&mut self) { + self.language_model.clear_cache(); + } +} diff --git a/src/models/w2v_bert_2_0/model.rs b/src/models/w2v_bert_2_0/model.rs index fb29025..7ff7972 100644 --- a/src/models/w2v_bert_2_0/model.rs +++ b/src/models/w2v_bert_2_0/model.rs @@ -7,7 +7,10 @@ use candle_nn::{ use crate::{ models::{ - common::{GLU, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm}, + common::{ + GLU, TwoLinearMLP, conv1d_depthwise, eager_attention_forward, get_conv1d, + get_layer_norm, + }, w2v_bert_2_0::config::W2VBert2_0Config, }, position_embed::rope::{RoPE, apply_rotary_pos_emb}, @@ -309,7 +312,12 @@ impl Wav2Vec2BertConvolutionModule { // (batch, channel, dim) let xs = self.glu.forward(&xs)?; let xs = xs.pad_with_zeros(D::Minus1, self.conv_depthwise_kernel_size - 1, 0)?; - let xs = self.depthwise_conv.forward(&xs)?; + // let xs = self.depthwise_conv.forward(&xs)?; + let xs = conv1d_depthwise( + &xs, + self.depthwise_conv.weight(), + self.depthwise_conv.bias(), + )?; let xs = self .depthwise_layer_norm .forward(&xs.transpose(1, 2)?)? diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index 319dfdb..c5013fc 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -142,10 +142,11 @@ pub fn glm_asr_apply_rotary_pos_emb( let cos = cos.to_dtype(q.dtype())?; let sin = sin.to_dtype(q.dtype())?; let rotary_dim = cos.dim(D::Minus1)?; + let q_dim = q.dim(D::Minus1)?; let q_rot = q.narrow(D::Minus1, 0, rotary_dim)?; - let q_pass = q.narrow(D::Minus1, rotary_dim, rotary_dim)?; + let q_pass = q.narrow(D::Minus1, rotary_dim, q_dim - rotary_dim)?; let k_rot = k.narrow(D::Minus1, 0, rotary_dim)?; - let k_pass = k.narrow(D::Minus1, rotary_dim, rotary_dim)?; + let k_pass = k.narrow(D::Minus1, rotary_dim, q_dim - rotary_dim)?; let q_embed = q_rot .broadcast_mul(&cos)? diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index a0bb812..fbf2953 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -554,7 +554,7 @@ pub fn l2_normalize(t: &Tensor, dim: usize) -> Result { if dim >= rank { return Err(anyhow!(format!("input dim {} must < rank {}", dim, rank))); } - let l2_norm = t.sqr()?.sum_keepdim(dim)?.sqrt()?; + let l2_norm = t.sqr()?.sum_keepdim(dim)?.affine(1.0, 1e-6)?.sqrt()?; Ok(t.broadcast_div(&l2_norm)?) } @@ -650,3 +650,29 @@ pub fn cosine_similarity(query_vector: &Tensor, matrix: &Tensor) -> Result Result { + if repeats == 1 { + return Ok(t.clone()); + } + let rank = t.rank(); + if dim >= rank { + return Err(anyhow!( + "Dimension {} is out of range for tensor with {} dimensions", + dim, + rank + )); + } + + let dims = t.dims(); + let mut indices = Vec::with_capacity(dims[dim] * repeats); + for i in 0..dims[dim] { + for _ in 0..repeats { + indices.push(i as u32); + } + } + + let indices_tensor = Tensor::from_vec(indices, (dims[dim] * repeats,), t.device())?; + let t = t.index_select(&indices_tensor, dim)?; + Ok(t) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index b6faf89..8b8d01a 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -5,7 +5,7 @@ // use std::io::{Read, Seek}; // use std::{io::Cursor, time::Instant}; -use aha::utils::interpolate::interpolate_nearest_2d; +// use aha::utils::tensor_utils::repeat_interleave; // use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::Result; // use byteorder::{LittleEndian, ReadBytesExt}; @@ -15,12 +15,23 @@ use candle_core::Tensor; #[test] fn messy_test() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda messy_test -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test messy_test messy_test -r -- --nocapture let device = &candle_core::Device::Cpu; - let input = Tensor::arange(0.0f32, 25.0f32, device)?.reshape((1, 1, 5, 5))?; - println!("input: {}", input); - let x_nearest = interpolate_nearest_2d(&input, (10, 10))?; - println!("x_nearest: {}", x_nearest); + let t1 = Tensor::randn(0.0, 1.0, (16, 9, 64, 128), device)?; + let t2 = Tensor::randn(0.0, 1.0, (16, 9, 128, 64), device)?; + let out = t1.matmul(&t2)?; + println!("out shape: {:?}", out); + + // let input = Tensor::arange(0.0f32, 25.0f32, device)?.reshape((5, 5))?; + // println!("input: {}", input); + // // let input = input.unsqueeze(D::Minus1)?; + // // let input = input.repeat((1, 1, 2))?; + // // let input = input.flatten(D::Minus2, D::Minus1)?; + // let output = repeat_interleave(&input, 2, 1)?; + // println!("output: {}", output); + + // let x_nearest = interpolate_nearest_2d(&input, (10, 10))?; + // println!("x_nearest: {}", x_nearest); // let input = Tensor::arange(0.0f32, 25.0f32, device)?.reshape((1, 5, 5))?; // println!("input: {}", input); // let x_nearest = interpolate_nearest_1d(&input, 10)?; diff --git a/tests/test_fun_asr_nano.rs b/tests/test_fun_asr_nano.rs index 552917c..a9af346 100644 --- a/tests/test_fun_asr_nano.rs +++ b/tests/test_fun_asr_nano.rs @@ -6,7 +6,7 @@ use anyhow::Result; use rocket::futures::StreamExt; #[test] fn fun_asr_nano_generate() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda fun_asr_nano_generate -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test test_fun_asr_nano fun_asr_nano_generate -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; let model_path = format!("{}/FunAudioLLM/Fun-ASR-Nano-2512/", save_dir); diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs index 78238d1..8e75549 100644 --- a/tests/test_glm_asr_nano.rs +++ b/tests/test_glm_asr_nano.rs @@ -7,7 +7,7 @@ use rocket::futures::StreamExt; #[test] fn glm_asr_nano_generate() -> Result<()> { - // RUST_BACKTRACE=1 cargo test -F cuda glm_asr_nano_generate -r -- --nocapture + // RUST_BACKTRACE=1 cargo test -F cuda --test test_glm_asr_nano glm_asr_nano_generate -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; let model_path = format!("{}/ZhipuAI/GLM-ASR-Nano-2512/", save_dir); diff --git a/tests/test_qwen3_5.rs b/tests/test_qwen3_5.rs new file mode 100644 index 0000000..567899a --- /dev/null +++ b/tests/test_qwen3_5.rs @@ -0,0 +1,105 @@ +use std::{pin::pin, time::Instant}; + +use aha::models::{GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; +use rocket::futures::StreamExt; + +#[test] +fn qwen3_5_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3_5 qwen3_5_generate -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/Qwen/Qwen3.5-0.8B/", save_dir); + + let message = r#" + { + "model": "qwen3.5", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image", + "image_url": + { + "url": "file:///home/jhq/Downloads/gougou1.jpg" + } + }, + { + "type": "text", + "text": "描述这张图片." + } + ] + } + ], + "metadata": {"enable_thinking": "true"} + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut qwen3vl = Qwen3_5GenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let res = qwen3vl.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("generate: \n {:?}", res); + if res.usage.is_some() { + let num_token = res.usage.as_ref().unwrap().total_tokens; + let duration_secs = i_duration.as_secs_f64(); + let tps = num_token as f64 / duration_secs; + println!("Tokens per second (TPS): {:.2}", tps); + } + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} + +#[tokio::test] +async fn qwen3_5_stream() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3_5 qwen3_5_stream -r -- --nocapture + + let save_dir = + aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; + let model_path = format!("{}/Qwen/Qwen3.5-0.8B/", save_dir); + + let message = r#" + { + "model": "qwen3.5", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image", + "image_url": + { + "url": "file:///home/jhq/Downloads/gougou1.jpg" + } + }, + { + "type": "text", + "text": "描述这张图片." + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let mut stream = pin!(qwen3_5.generate_stream(mes)?); + while let Some(item) = stream.next().await { + println!("generate: \n {:?}", item); + } + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index 7f0f733..031213b 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -59,7 +59,7 @@ fn qwen3vl_thinking_generate() -> Result<()> { #[test] fn qwen3vl_generate() -> Result<()> { - // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg qwen3vl_generate -r -- --nocapture + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda --test test_qwen3vl qwen3vl_generate -r -- --nocapture let save_dir = aha::utils::get_default_save_dir().ok_or(anyhow::anyhow!("Failed to get save dir"))?; @@ -73,15 +73,15 @@ fn qwen3vl_generate() -> Result<()> { "role": "user", "content": [ { - "type": "video", - "video_url": + "type": "image", + "image_url": { - "url": "./assets/video/video_test.mp4" + "url": "file:///home/jhq/Downloads/gougou1.jpg" } - }, + }, { "type": "text", - "text": "视频中发生了什么?" + "text": "描述这张图片." } ] }