add Qwen3.5 model

This commit is contained in:
jhqxxx
2026-03-05 21:57:03 +08:00
parent a61b899741
commit 4ae8b0b2f4
31 changed files with 1930 additions and 121 deletions
Generated
+1 -1
View File
@@ -30,7 +30,7 @@ dependencies = [
[[package]]
name = "aha"
version = "0.2.0"
version = "0.2.1"
dependencies = [
"aha_openai_dive",
"anyhow",
+2 -2
View File
@@ -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" }
+3
View File
@@ -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
+8
View File
@@ -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 语音识别模型
+9
View File
@@ -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
+9
View File
@@ -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
### 新增
+24 -21
View File
@@ -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
+22 -19
View File
@@ -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) 获取最新版本。
模型定期更新。
## 性能基准
+8
View File
@@ -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",
+21
View File
@@ -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<bool>,
) -> Result<String> {
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)
}
}
+1
View File
@@ -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;
+103
View File
@@ -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(())
}
}
+1 -1
View File
@@ -61,7 +61,7 @@ impl ExecModel for Qwen3vlExec {
} else {
format!(
r#"{{
"model": "qwen2.5vl",
"model": "qwen3vl",
"messages": [
{{
"role": "user",
+2
View File
@@ -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};
+20
View File
@@ -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)?;
+43 -45
View File
@@ -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<Tensor> {
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::<Result<Vec<Tensor>>>()?;
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<Tensor> {
let xs = xs.mul(&tanh)?;
Ok(xs)
}
pub fn softplus(xs: &Tensor) -> Result<Tensor> {
// ln(1 + exp(x))
Ok((xs.exp()? + 1.0)?.log()?)
}
pub fn softplus_stable(xs: &Tensor) -> Result<Tensor> {
// 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<Tensor> {
// 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)?)
}
}
}
+4 -2
View File
@@ -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 {
+5 -4
View File
@@ -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<Tensor> {
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()?;
+42 -7
View File
@@ -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<ModelInstance<'_
let model = Qwen3GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3(model)
}
WhichModel::Qwen3_5_0_8B => {
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)
+57
View File
@@ -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<usize>,
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<String>,
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<usize>,
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,
}
+231
View File
@@ -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<DType>) -> Result<Self> {
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<ChatCompletionResponse> {
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let enable_thinking = extract_metadata_value::<bool>(&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<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ 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::<bool>(&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>" => {
// 开始工具调用
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;
}
"</tool_call>" => {
// 结束工具调用
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)))
}
}
+3
View File
@@ -0,0 +1,3 @@
pub mod config;
pub mod generate;
pub mod model;
File diff suppressed because it is too large Load Diff
+10 -2
View File
@@ -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)?)?
+3 -2
View File
@@ -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)?
+27 -1
View File
@@ -554,7 +554,7 @@ pub fn l2_normalize(t: &Tensor, dim: usize) -> Result<Tensor> {
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<Tenso
.squeeze(D::Minus1)?;
Ok(similarity)
}
pub fn repeat_interleave(t: &Tensor, repeats: usize, dim: usize) -> Result<Tensor> {
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)
}
+17 -6
View File
@@ -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)?;
+1 -1
View File
@@ -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);
+1 -1
View File
@@ -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);
+105
View File
@@ -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(())
}
+6 -6
View File
@@ -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": "描述这张图片."
}
]
}