merge main

This commit is contained in:
jhqxxx
2026-02-02 17:56:34 +08:00
28 changed files with 2881 additions and 84 deletions
+108
View File
@@ -0,0 +1,108 @@
# Changelog
## [Unreleased]
### Added
- **CLI `run` Subcommand**: Direct model inference from CLI without HTTP service overhead:
- `aha run` - Run model inference directly
- `-m, --model <MODEL>` - Specify which model to use
- `-in, --input <INPUT>` - Input text or file path (model-specific interpretation)
- `-out, --output <OUTPUT>` - Output file path (optional, auto-generated if not specified)
- `--weight-path <WEIGHT_PATH>` - Local model weight path (required)
### Changed
- **CLI Structure**: Refactored CLI to use clap's Subcommand feature while maintaining backward compatibility
- **Backward Compatibility**: Commands without subcommand now default to `cli` subcommand:
- `aha -m qwen3vl-2b` is equivalent to `aha cli -m qwen3vl-2b`
- All existing parameter options and defaults remain unchanged
### Technical Details
**Subcommand Parameters:**
`aha cli`:
- `-a, --address <ADDRESS>` - Server address (default: 127.0.0.1)
- `-p, --port <PORT>` - Server port (default: 10100)
- `-m, --model <MODEL>` - Model to use (required)
- `--weight-path <WEIGHT_PATH>` - Local model weight path (optional)
- `--save-dir <SAVE_DIR>` - Directory to save downloaded model (optional)
- `--download-retries <DOWNLOAD_RETRIES>` - Download retry attempts (default: 3)
`aha serv`:
- `-a, --address <ADDRESS>` - Server address (default: 127.0.0.1)
- `-p, --port <PORT>` - Server port (default: 10100)
- `-m, --model <MODEL>` - Model to use (required)
- `--weight-path <WEIGHT_PATH>` - Local model weight path (required)
`aha download`:
- `-m, --model <MODEL>` - Model to download (required)
- `-s, --save-dir <SAVE_DIR>` - Directory to save downloaded model (optional)
- `--download-retries <DOWNLOAD_RETRIES>` - Download retry attempts (default: 3)
**Code Changes:**
- Modified `src/main.rs` only
- Extracted common functions: `get_model_id()`, `start_http_server()`
- Reused existing `download_model()` and `init()` functions
- No changes to other modules or dependencies
## [0.1.8] - 2025-01-20
### Added
- Support for Fun-ASR-Nano-2512 model
- Support for Qwen3-0.6B model
## [0.1.7] - 2024-XX-XX
### Added
- Support for GLM-ASR-Nano-2512 model
## [0.1.6] - 2024-XX-XX
### Added
- Support for RMBG-2.0 model (background removal)
## [0.1.5] - 2024-XX-XX
### Added
- Support for VoxCPM1.5 model
## [0.1.4] - 2024-XX-XX
### Added
- Support for PaddleOCR-VL model
## [0.1.3] - 2024-XX-XX
### Added
- Support for Hunyuan-OCR model
## [0.1.2] - 2024-XX-XX
### Added
- Support for DeepSeek-OCR model
## [0.1.1] - 2024-XX-XX
### Added
- Support for Qwen3VL model family (2B, 4B, 8B, 32B)
## [0.1.0] - 2024-XX-XX
### Added
- Initial release
- Support for Qwen2.5VL models (3B, 7B)
- Support for MiniCPM4-0.5B model
- Support for VoxCPM-0.5B model
+4
View File
@@ -10,6 +10,10 @@ build:
@echo "Building project..."
@cargo build
build_mac:
@echo "Building project for macOS..."
@cargo build --features metal --release
test:
@echo "Running tests..."
@cargo test
+291
View File
@@ -0,0 +1,291 @@
# AHA 命令行使用说明
## 概述
AHA 是一个基于 Candle 框架的高性能模型推理库,支持多种多模态模型,包括视觉、语言和语音模型。
```bash
aha [COMMAND] [OPTIONS]
```
## 全局选项
| 选项 | 说明 | 默认值 |
|------|------|--------|
| `-a, --address <ADDRESS>` | 服务监听地址 | 127.0.0.1 |
| `-p, --port <PORT>` | 服务监听端口 | 10100 |
| `-m, --model <MODEL>` | 模型类型(必选) | - |
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径 | - |
| `--save-dir <SAVE_DIR>` | 模型下载保存目录 | ~/.aha/ |
| `--download-retries <DOWNLOAD_RETRIES>` | 下载重试次数 | 3 |
| `-h, --help` | 显示帮助信息 | - |
| `-V, --version` | 显示版本号 | - |
## 子命令
### cli - 下载模型并启动服务(默认)
下载指定的模型并启动 HTTP 服务。当不指定子命令时,默认使用此命令。
**语法:**
```bash
aha cli [OPTIONS] --model <MODEL>
```
**选项:**
| 选项 | 说明 | 默认值 |
|------|------|--------|
| `-a, --address <ADDRESS>` | 服务监听地址 | 127.0.0.1 |
| `-p, --port <PORT>` | 服务监听端口 | 10100 |
| `-m, --model <MODEL>` | 模型类型(必选) | - |
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径(如指定则跳过下载) | - |
| `--save-dir <SAVE_DIR>` | 模型下载保存目录 | ~/.aha/ |
| `--download-retries <DOWNLOAD_RETRIES>` | 下载重试次数 | 3 |
**示例:**
```bash
# 下载模型并启动服务(默认端口 10100)
aha cli -m qwen3vl-2b
# 指定端口和保存目录
aha cli -m qwen3vl-2b -p 8080 --save-dir /data/models
# 使用本地模型(不下载)
aha cli -m qwen3vl-2b --weight-path /path/to/model
# 向后兼容方式(等同于 cli 子命令)
aha -m qwen3vl-2b
```
### run - 直接模型推理
直接运行模型推理,无需启动 HTTP 服务。适用于一次性推理任务或批处理。
**语法:**
```bash
aha run [OPTIONS] --model <MODEL> --input <INPUT> [--input <INPUT2>] --weight-path <WEIGHT_PATH>
```
**选项:**
| 选项 | 说明 | 默认值 |
|------|------|--------|
| `-m, --model <MODEL>` | 模型类型(必选) | - |
| `-i, --input <INPUT>` | 输入文本或文件路径(模型特定解释,支持1-2个参数, input1: 提示文本, input2: 文件地址) | - |
| `-o, --output <OUTPUT>` | 输出文件路径(可选,未指定则自动生成) | - |
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径(必选) | - |
**示例:**
```bash
# VoxCPM1.5 文字转语音(单个输入)
aha run -m voxcpm1.5 -i "太阳当空照" -o output.wav --weight-path /path/to/model
# VoxCPM1.5 从文件读取输入(单个输入)
aha run -m voxcpm1.5 -i "file://./input.txt" --weight-path /path/to/model
# MiniCPM4 文本生成(单个输入)
aha run -m minicpm4-0.5b -i "你好" --weight-path /path/to/model
# DeepSeek OCR 图片识别(单个输入)
aha run -m deepseek-ocr -i "image.jpg" --weight-path /path/to/model
# RMBG2.0 背景移除(单个输入)
aha run -m RMBG2.0 -i "photo.png" -o "no_bg.png" --weight-path /path/to/model
# GLM-ASR 语音识别(两个输入:提示文本 + 音频文件)
aha run -m glm-asr-nano-2512 -i "请转写这段音频" -i "audio.wav" --weight-path /path/to/model
# Fun-ASR 语音识别(两个输入:提示文本 + 音频文件)
aha run -m fun-asr-nano-2512 -i "语音转写:" -i "audio.wav" --weight-path /path/to/model
# qwen3 文本生成(单个输入)
aha run -m qwen3-0.6b -i "你好" --weight-path /path/to/model
# qwen2.5vl 图像理解(两个输入:提示文本 + 图片文件)
aha run -m qwen2.5vl-3b -i "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本" -i "image.jpg" --weight-path /path/to/model
```
### serv - 启动服务
仅启动 HTTP 服务,不下载模型。必须通过 `--weight-path` 指定本地模型路径。
**语法:**
```bash
aha serv [OPTIONS] --model <MODEL> --weight-path <WEIGHT_PATH>
```
**选项:**
| 选项 | 说明 | 默认值 |
|------|------|--------|
| `-a, --address <ADDRESS>` | 服务监听地址 | 127.0.0.1 |
| `-p, --port <PORT>` | 服务监听端口 | 10100 |
| `-m, --model <MODEL>` | 模型类型(必选) | - |
| `--weight-path <WEIGHT_PATH>` | 本地模型权重路径(必选) | - |
**示例:**
```bash
# 使用本地模型启动服务
aha serv -m qwen3vl-2b --weight-path /path/to/model
# 指定端口启动
aha serv -m qwen3vl-2b --weight-path /path/to/model -p 8080
# 指定监听地址
aha serv -m qwen3vl-2b --weight-path /path/to/model -a 0.0.0.0
```
### download - 下载模型
仅下载指定模型,不启动服务。
**语法:**
```bash
aha download [OPTIONS] --model <MODEL>
```
**选项:**
| 选项 | 说明 | 默认值 |
|------|------|--------|
| `-m, --model <MODEL>` | 模型类型(必选) | - |
| `-s, --save-dir <SAVE_DIR>` | 模型下载保存目录 | ~/.aha/ |
| `--download-retries <DOWNLOAD_RETRIES>` | 下载重试次数 | 3 |
**示例:**
```bash
# 下载模型到默认目录
aha download -m qwen3vl-2b
# 指定保存目录
aha download -m qwen3vl-2b -s /data/models
# 指定下载重试次数
aha download -m qwen3vl-2b --download-retries 5
# 下载 MiniCPM4-0.5B 模型
aha download -m minicpm4-0.5b -s models
```
## 支持的模型
| 模型标识 | 模型名称 | 说明 |
|---------|---------|------|
| `minicpm4-0.5b` | OpenBMB/MiniCPM4-0.5B | 面壁智能 MiniCPM4 0.5B 模型 |
| `qwen2.5vl-3b` | Qwen/Qwen2.5-VL-3B-Instruct | 通义千问 2.5 VL 3B 模型 |
| `qwen2.5vl-7b` | Qwen/Qwen2.5-VL-7B-Instruct | 通义千问 2.5 VL 7B 模型 |
| `qwen3-0.6b` | Qwen/Qwen3-0.6B | 通义千问 3 0.6B 模型 |
| `qwen3vl-2b` | Qwen/Qwen3-VL-2B-Instruct | 通义千问 3 VL 2B 模型 |
| `qwen3vl-4b` | Qwen/Qwen3-VL-4B-Instruct | 通义千问 3 VL 4B 模型 |
| `qwen3vl-8b` | Qwen/Qwen3-VL-8B-Instruct | 通义千问 3 VL 8B 模型 |
| `qwen3vl-32b` | Qwen/Qwen3-VL-32B-Instruct | 通义千问 3 VL 32B 模型 |
| `deepseek-ocr` | deepseek-ai/DeepSeek-OCR | DeepSeek OCR 模型 |
| `hunyuan-ocr` | Tencent-Hunyuan/HunyuanOCR | 腾讯混元 OCR 模型 |
| `paddleocr-vl` | PaddlePaddle/PaddleOCR-VL | 百度飞桨 OCR VL 模型 |
| `RMBG2.0` | AI-ModelScope/RMBG-2.0 | RMBG 2.0 背景移除模型 |
| `voxcpm` | OpenBMB/VoxCPM-0.5B | 面壁智能 VoxCPM 0.5B 语音生成模型 |
| `voxcpm1.5` | OpenBMB/VoxCPM1.5 | 面壁智能 VoxCPM 1.5 语音生成模型 |
| `glm-asr-nano-2512` | ZhipuAI/GLM-ASR-Nano-2512 | 智谱 AI ASR Nano 2512 语音识别模型 |
| `fun-asr-nano-2512` | FunAudioLLM/Fun-ASR-Nano-2512 | 通义百聆 ASR Nano 2512 语音识别模型 |
## 常见使用场景
### 场景 1:快速启动推理服务
```bash
# 一条命令下载并启动服务
aha -m qwen3vl-2b
```
### 场景 2:使用已有模型启动服务
```bash
# 假设模型已下载到 /data/models/Qwen/Qwen3-VL-2B-Instruct
aha serv -m qwen3vl-2b --weight-path /data/models/Qwen/Qwen3-VL-2B-Instruct
```
### 场景 3:预先下载模型
```bash
# 下载模型到指定目录,稍后使用
aha download -m qwen3vl-2b -s /data/models
# 后续启动时直接使用
aha serv -m qwen3vl-2b --weight-path /data/models/Qwen/Qwen3-VL-2B-Instruct
```
### 场景 4:自定义服务端口和地址
```bash
# 在 0.0.0.0:8080 启动服务,允许外部访问
aha -m qwen3vl-2b -a 0.0.0.0 -p 8080
```
## API 接口
服务启动后,提供以下 API 接口:
### 对话接口
- **端点**: `POST /chat/completions`
- **功能**: 多模态对话和文本生成
- **支持模型**: Qwen2.5VL, Qwen3, Qwen3VL, DeepSeekOCR, GLM-ASR-Nano-2512, Fun-ASR-Nano-2512 等
- **格式**: OpenAI Chat Completion 格式
- **流式支持**: 支持
### 图像处理接口
- **端点**: `POST /images/remove_background`
- **功能**: 图像背景移除
- **支持模型**: RMBG-2.0
- **格式**: OpenAI Chat Completion 格式
- **流式支持**: 不支持
### 语音生成接口
- **端点**: `POST /audio/speech`
- **功能**: 语音合成和生成
- **支持模型**: VoxCPM, VoxCPM1.5
- **格式**: OpenAI Chat Completion 格式
- **流式支持**: 不支持
## 向后兼容性
为了保持与旧版本的兼容性,以下两种使用方式是等效的:
```bash
# 新方式(推荐)
aha cli -m qwen3vl-2b
# 旧方式(向后兼容)
aha -m qwen3vl-2b
```
## 注意事项
1. **serv 子命令必须指定 `--weight-path`**:由于 `serv` 子命令不下载模型,必须通过 `--weight-path` 指定已下载的模型路径。
2. **下载重试机制**:默认重试 3 次,每次失败后等待 2 秒再重试。可通过 `--download-retries` 调整重试次数。
3. **默认保存目录**:模型默认保存到 `~/.aha/` 目录下,可通过 `--save-dir``-d` 参数自定义。
4. **端口占用**:启动服务前确保指定的端口未被占用,默认端口为 10100。
5. **权限问题**:如果保存到系统目录(如 `/data/models`),确保有相应的写入权限。
## 获取帮助
```bash
# 查看主帮助
aha --help
# 查看子命令帮助
aha cli --help
aha serv --help
aha download --help
# 查看版本信息
aha --version
```
+4 -4
View File
@@ -19,12 +19,12 @@ pub fn get_template(path: String) -> Result<String> {
// 修复模板中的问题行
let fixed_template = chat_template
.replace(
"message.content.startswith('<tool_response>')",
"message.content is startingwith('<tool_response>')", // 使用minijinja中的 is startingwith 替换
"content.startswith('<tool_response>')",
"content is startingwith('<tool_response>')", // 使用minijinja中的 is startingwith 替换
)
.replace(
"message.content.endswith('</tool_response>')",
"message.content is endingwith('</tool_response>')", // 使用minijinja中的 is endingwith 替换
"content.endswith('</tool_response>')",
"content is endingwith('</tool_response>')", // 使用minijinja中的 is endingwith 替换
)
.replace(
"content.split('</think>')[0].rstrip('\\n').split('<think>')[-1].lstrip('\\n')",
+68
View File
@@ -0,0 +1,68 @@
//! DeepSeek-OCR exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, deepseek_ocr::generate::DeepseekOCRGenerateModel};
pub struct DeepSeekORExec;
impl ExecModel for DeepSeekORExec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
let url = &input[0];
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
{
url.clone()
} else {
format!("file://{}", url)
};
let i_start = Instant::now();
let mut model = DeepseekOCRGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"model": "deepseek-ocr",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "image",
"image_url": {{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "<image>\nConvert the document to markdown. "
}}
]
}}
]
}}"#,
input_url
);
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(())
}
}
+79
View File
@@ -0,0 +1,79 @@
//! Fun-ASR-Nano-2512 exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, fun_asr_nano::generate::FunAsrNanoGenerateModel};
use crate::utils::get_file_path;
pub struct FunASRNanoExec;
impl ExecModel for FunASRNanoExec {
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 = &input[7..];
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 = FunAsrNanoGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
// Create ChatCompletionParameters for ASR
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 = format!(
r#"{{
"model": "fun-asr-nano",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "audio",
"audio_url": {{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "{}"
}}
]
}}
]
}}"#,
input_url, target_text
);
let mes = serde_json::from_str(&message)?;
let i_start = Instant::now();
let res = model.generate(mes)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
println!("Result: {:?}", res);
if let Some(out) = output {
std::fs::write(out, format!("{:?}", res))?;
println!("Output saved to: {}", out);
}
Ok(())
}
}
+80
View File
@@ -0,0 +1,80 @@
//! GLM-ASR-Nano-2512 exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, glm_asr_nano::generate::GlmAsrNanoGenerateModel};
use crate::utils::get_file_path;
pub struct GlmASRNanoExec;
impl ExecModel for GlmASRNanoExec {
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 = &input[7..];
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 = GlmAsrNanoGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
// Create ChatCompletionParameters for ASR
// Input should be an audio file path
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 = format!(
r#"{{
"model": "glm-asr-nano",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "audio",
"audio_url": {{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "{}"
}}
]
}}
]
}}"#,
input_url, target_text
);
let mes = serde_json::from_str(&message)?;
let i_start = Instant::now();
let res = model.generate(mes)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
println!("Result: {:?}", res);
if let Some(out) = output {
std::fs::write(out, format!("{:?}", res))?;
println!("Output saved to: {}", out);
}
Ok(())
}
}
+68
View File
@@ -0,0 +1,68 @@
//! Hunyuan-OCR exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, hunyuan_ocr::generate::HunyuanOCRGenerateModel};
pub struct HunyuanORExec;
impl ExecModel for HunyuanORExec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
let url = &input[0];
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
{
url.clone()
} else {
format!("file://{}", url)
};
let i_start = Instant::now();
let mut model = HunyuanOCRGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"model": "hunyuan-ocr",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "image",
"image_url": {{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "检测并识别图片中的文字,将文本坐标格式化输出。"
}}
]
}}
]
}}"#,
input_url
);
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(())
}
}
+60
View File
@@ -0,0 +1,60 @@
//! MiniCPM4-0.5B exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, minicpm4::generate::MiniCPMGenerateModel};
use crate::utils::get_file_path;
pub struct MiniCPM4Exec;
impl ExecModel for MiniCPM4Exec {
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 = &input[7..];
let path = get_file_path(input_text)?;
std::fs::read_to_string(path)?
} else {
input_text.to_string()
};
let i_start = Instant::now();
let mut model = MiniCPMGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"temperature": 0.3,
"top_p": 0.8,
"model": "minicpm4",
"messages": [
{{
"role": "user",
"content": "{}"
}}
]
}}"#,
target_text.replace('"', "\\\"")
);
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);
// Print result
println!("Result: {:?}", result);
if let Some(out) = output {
std::fs::write(out, format!("{:?}", result))?;
println!("Output saved to: {}", out);
}
Ok(())
}
}
+37
View File
@@ -0,0 +1,37 @@
//! CLI exec module for direct model inference
//!
//! This module provides model-specific exec implementations for the `run` subcommand.
//! Each model has its own exec module that handles input/output parsing and model invocation.
pub mod deepseek_ocr;
pub mod fun_asr_nano;
pub mod glm_asr_nano;
pub mod hunyuan_ocr;
pub mod minicpm4;
pub mod paddleocr_vl;
pub mod qwen2_5vl;
pub mod qwen3;
pub mod qwen3vl;
pub mod rmbg2_0;
pub mod voxcpm;
pub mod voxcpm1_5;
use anyhow::Result;
/// Trait for model exec implementations
///
/// Each model exec module implements this trait to provide
/// model-specific inference logic for CLI `run` commands.
pub trait ExecModel {
/// Run inference with the given input and output parameters
///
/// # Arguments
/// * `input` - Input text or file path (interpretation is model-specific)
/// * `output` - Optional output file path (if None, model will auto-generate)
/// * `weight_path` - Path to the model weights
///
/// # Returns
/// * `Ok(())` on success
/// * `Err(anyhow::Error)` on failure
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()>;
}
+68
View File
@@ -0,0 +1,68 @@
//! PaddleOCR-VL exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, paddleocr_vl::generate::PaddleOCRVLGenerateModel};
pub struct PaddleOVLExec;
impl ExecModel for PaddleOVLExec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
let url = &input[0];
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
{
url.clone()
} else {
format!("file://{}", url)
};
let i_start = Instant::now();
let mut model = PaddleOCRVLGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"model": "paddleocr-vl",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "image",
"image_url": {{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "OCR:"
}}
]
}}
]
}}"#,
input_url
);
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(())
}
}
+75
View File
@@ -0,0 +1,75 @@
//! Qwen2.5VL-3B exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel};
use crate::utils::get_file_path;
pub struct Qwen2_5vlExec;
impl ExecModel for Qwen2_5vlExec {
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 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 i_start = Instant::now();
let mut model = Qwen2_5VLGenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"model": "qwen2.5vl",
"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(())
}
}
+57
View File
@@ -0,0 +1,57 @@
//! Qwen3-0.6B exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, qwen3::generate::Qwen3GenerateModel};
use crate::utils::get_file_path;
pub struct Qwen3Exec;
impl ExecModel for Qwen3Exec {
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 = &input[7..];
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 = Qwen3GenerateModel::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!(
r#"{{
"model": "qwen3",
"messages": [
{{
"role": "user",
"content": "{}"
}}
]
}}"#,
target_text.replace('"', "\\\"")
);
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(())
}
}
+102
View File
@@ -0,0 +1,102 @@
//! Qwen3VL-2B exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::{GenerateModel, qwen3vl::generate::Qwen3VLGenerateModel};
use crate::utils::get_file_path;
pub struct Qwen3vlExec;
impl ExecModel for Qwen3vlExec {
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 = Qwen3VLGenerateModel::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": "qwen3vl",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "video",
"video_url":
{{
"url": "{}"
}}
}},
{{
"type": "text",
"text": "{}"
}}
]
}}
]
}}"#,
input_url, target_text
)
} else {
format!(
r#"{{
"model": "qwen2.5vl",
"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(())
}
}
+78
View File
@@ -0,0 +1,78 @@
//! RMBG2.0 exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::exec::ExecModel;
use crate::models::rmbg2_0::generate::RMBG2_0Model;
pub struct RMBG2_0Exec;
impl ExecModel for RMBG2_0Exec {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
let url = &input[0];
let input_url = if url.starts_with("http://")
|| url.starts_with("https://")
|| url.starts_with("file://")
{
url.clone()
} else {
format!("file://{}", url)
};
let i_start = Instant::now();
let model = RMBG2_0Model::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
// Create ChatCompletionParameters for image background removal
let message = format!(
r#"{{
"model": "rmbg2.0",
"messages": [
{{
"role": "user",
"content": [
{{
"type": "image",
"image_url": {{
"url": "{}"
}}
}}
]
}}
]
}}"#,
input_url
);
let mes = serde_json::from_str(&message)?;
let i_start = Instant::now();
let result = model.inference(mes)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
let output_path = if let Some(out) = output {
out.to_string()
} else {
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)?
.as_secs();
format!("rmbg_{}.png", timestamp)
};
// Save all result images
for (i, img) in result.iter().enumerate() {
let path = if result.len() == 1 {
output_path.clone()
} else {
format!("{}_{}.png", output_path.trim_end_matches(".png"), i)
};
img.save(&path)?;
println!("Output saved to: {}", path);
}
Ok(())
}
}
+58
View File
@@ -0,0 +1,58 @@
//! VoxCPM exec implementation for CLI `run` subcommand
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::models::voxcpm::generate::VoxCPMGenerate;
use crate::{exec::ExecModel, utils::get_file_path};
pub struct VoxCPMExec;
impl ExecModel for VoxCPMExec {
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 = &input[7..];
let path = get_file_path(input_text)?;
std::fs::read_to_string(path)?
} else {
input_text.clone()
};
let i_start = Instant::now();
let mut voxcpm_generate = VoxCPMGenerate::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now();
let audio = voxcpm_generate.inference(
target_text,
Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()), // todo args
Some("file://./assets/audio/voice_01.wav".to_string()), // todo args
2,
100, // max_len (voxcpm uses 100 vs voxcpm1.5's 4096)
10,
2.0,
6.0,
)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
let output_path = if let Some(out) = output {
out.to_string()
} else {
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)?
.as_secs();
format!("voxcpm_{}.wav", timestamp)
};
let sample_rate = voxcpm_generate.sample_rate();
crate::utils::audio_utils::save_wav(&audio, &output_path, sample_rate as u32)?;
println!("Output saved to: {}", output_path);
Ok(())
}
}
+63
View File
@@ -0,0 +1,63 @@
//! VoxCPM1.5 exec implementation for CLI `run` subcommand
//!
//! This module handles VoxCPM1.5 model inference for direct CLI execution.
//! Input/output parameter interpretation is handled here as per the design:
//! - Input can be text content or a file path (with `file://` prefix)
//! - Output can be a file path or will be auto-generated if not specified
use std::time::Instant;
use anyhow::{Ok, Result};
use crate::models::voxcpm::generate::VoxCPMGenerate;
use crate::{exec::ExecModel, utils::get_file_path};
pub struct VoxCPM1_5Exec;
impl ExecModel for VoxCPM1_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 = &input[7..];
let path = get_file_path(input_text)?;
std::fs::read_to_string(path)?
} else {
input_text.clone()
};
let i_start = Instant::now();
let mut voxcpm_generate = VoxCPMGenerate::init(weight_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now();
let audio = voxcpm_generate.inference(
target_text,
Some("啥子小师叔,打狗还要看主人,你再要继续,我就是你的对手".to_string()), // todo args
Some("file://./assets/audio/voice_01.wav".to_string()), // todo args
2,
4096,
10,
2.0,
6.0,
)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in generate is: {:?}", i_duration);
let output_path = if let Some(out) = output {
out.to_string()
} else {
let timestamp = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)?
.as_secs();
format!("voxcpm1_5_{}.wav", timestamp)
};
let sample_rate = voxcpm_generate.sample_rate();
crate::utils::audio_utils::save_wav(&audio, &output_path, sample_rate as u32)?;
println!("Output saved to: {}", output_path);
Ok(())
}
}
+1
View File
@@ -1,4 +1,5 @@
pub mod chat_template;
pub mod exec;
pub mod models;
pub mod position_embed;
pub mod tokenizer;
+323 -64
View File
@@ -1,7 +1,8 @@
use std::{net::IpAddr, str::FromStr, time::Duration};
use aha::{models::WhichModel, utils::{download_model, get_default_save_dir}};
use clap::Parser;
use clap::{Args, Parser, Subcommand, ValueEnum};
use modelscope::ModelScope;
use rocket::{
Config,
@@ -14,63 +15,136 @@ use crate::api::init;
mod api;
#[derive(Parser, Debug)]
#[command(name = "aha")]
#[command(version, about, long_about = None)]
struct Args {
struct Cli {
/// Service listen address
#[arg(short, long, default_value = "127.0.0.1")]
address: String,
#[arg(short, long, default_value_t = 10100)]
port: u16,
address: Option<String>,
/// Service listen port
#[arg(short, long)]
model: WhichModel,
port: Option<u16>,
/// Model type (required for backward compatibility)
#[arg(short, long)]
model: Option<WhichModel>,
/// Local model weight path
#[arg(long)]
weight_path: Option<String>,
/// Model download save directory
#[arg(long)]
save_dir: Option<String>,
/// Download retry count
#[arg(long)]
download_retries: Option<u32>,
#[command(subcommand)]
command: Option<Commands>,
}
#[derive(Subcommand, Debug)]
enum Commands {
/// Download model and start service (default)
Cli(CliArgs),
/// Start service only (requires --weight-path)
Serv(ServArgs),
/// Download model only
Download(DownloadArgs),
/// Run model inference directly
Run(RunArgs),
/// List all supported models
List,
}
/// Common/shared arguments for server operations
#[derive(Args, Debug)]
struct CommonArgs {
/// Service listen address
#[arg(short, long, default_value = "127.0.0.1")]
address: String,
/// Service listen port
#[arg(short, long, default_value_t = 10100)]
port: u16,
/// Model type (required)
#[arg(short, long)]
model: WhichModel,
}
/// Arguments for the 'cli' subcommand (download + serve)
#[derive(Args, Debug)]
struct CliArgs {
#[command(flatten)]
common: CommonArgs,
/// Local model weight path (skip download if provided)
#[arg(long)]
weight_path: Option<String>,
/// Model download save directory
#[arg(long)]
save_dir: Option<String>,
/// Download retry count
#[arg(long)]
download_retries: Option<u32>,
}
// async fn download_model(model_id: &str, save_dir: &str, max_retries: u32) -> anyhow::Result<()> {
// let mut attempts = 0u32;
// loop {
// attempts += 1;
// println!(
// "Attempting to download model (attempt {}/{})",
// attempts, max_retries
// );
// match ModelScope::download(model_id, save_dir).await {
// Ok(()) => {
// println!("Model downloaded successfully");
// return Ok(());
// }
// Err(e) => {
// if attempts >= max_retries {
// return Err(anyhow::anyhow!(
// "Failed to download model after {} attempts. Last error: {}",
// max_retries,
// e
// ));
// }
/// Arguments for the 'serv' subcommand (serve only)
#[derive(Args, Debug)]
struct ServArgs {
#[command(flatten)]
common: CommonArgs,
// println!(
// "Download failed (attempt {}): {}. Retrying in 2 seconds...",
// attempts, e
// );
// sleep(Duration::from_secs(2)).await;
// }
// }
// }
// }
/// Local model weight path (required)
#[arg(long, required = true)]
weight_path: String,
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let args = Args::parse();
let model_id = match &args.model {
/// Arguments for the 'download' subcommand (download only)
#[derive(Args, Debug)]
struct DownloadArgs {
/// Model type (required)
#[arg(short, long)]
model: WhichModel,
/// Model download save directory
#[arg(short, long)]
save_dir: Option<String>,
/// Download retry count
#[arg(long)]
download_retries: Option<u32>,
}
/// Arguments for the 'run' subcommand (direct inference)
#[derive(Args, Debug)]
struct RunArgs {
/// Model type (required)
#[arg(short, long)]
model: WhichModel,
/// Input text or file path
#[arg(short, long, num_args = 1..=2, value_delimiter = ' ')]
input: Vec<String>,
/// Output file path (optional)
#[arg(short, long)]
output: Option<String>,
/// Local model weight path (required)
#[arg(long, required = true)]
weight_path: String,
}
/// Get the ModelScope model ID for a given WhichModel variant
fn get_model_id(model: WhichModel) -> &'static str {
match model {
WhichModel::MiniCPM4_0_5B => "OpenBMB/MiniCPM4-0.5B",
WhichModel::Qwen2_5vl3B => "Qwen/Qwen2.5-VL-3B-Instruct",
WhichModel::Qwen2_5vl7B => "Qwen/Qwen2.5-VL-7B-Instruct",
@@ -87,30 +161,219 @@ async fn main() -> anyhow::Result<()> {
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
WhichModel::GlmASRNano2512 => "ZhipuAI/GLM-ASR-Nano-2512",
WhichModel::FunASRNano2512 => "FunAudioLLM/Fun-ASR-Nano-2512",
};
let model_path = match &args.weight_path {
Some(path) => path.clone(),
None => {
let save_dir = match &args.save_dir {
Some(dir) => dir.clone(),
None => get_default_save_dir().expect("Failed to get home directory"),
};
let max_retries = args.download_retries.unwrap_or(3);
download_model(model_id, &save_dir, max_retries).await?;
save_dir + "/" + model_id
}
};
// println!("-------------------download path: {}", model_path);
init(args.model, model_path)?;
start_http_server(&args).await?;
}
}
/// List all supported models
fn run_list() -> anyhow::Result<()> {
let models = [
WhichModel::MiniCPM4_0_5B,
WhichModel::Qwen2_5vl3B,
WhichModel::Qwen2_5vl7B,
WhichModel::Qwen3_0_6B,
WhichModel::Qwen3vl2B,
WhichModel::Qwen3vl4B,
WhichModel::Qwen3vl8B,
WhichModel::Qwen3vl32B,
WhichModel::DeepSeekOCR,
WhichModel::HunyuanOCR,
WhichModel::PaddleOCRVL,
WhichModel::RMBG2_0,
WhichModel::VoxCPM,
WhichModel::VoxCPM1_5,
WhichModel::GlmASRNano2512,
WhichModel::FunASRNano2512,
];
println!("Available models:");
println!();
println!("{:<30} ModelScope ID", "Model Name");
println!("{}", "-".repeat(80));
for model in models {
let possible_value = model.to_possible_value().unwrap();
let name = possible_value.get_name();
let id = get_model_id(model);
println!("{:<30} {}", name, id);
}
Ok(())
}
pub(crate) async fn start_http_server(args: &Args) -> anyhow::Result<()> {
/// Run the 'cli' subcommand: download model (if needed) and start service
async fn run_cli(args: CliArgs) -> anyhow::Result<()> {
let CliArgs {
common,
weight_path,
save_dir,
download_retries,
} = args;
let model_id = get_model_id(common.model);
let model_path = match weight_path {
Some(path) => path,
None => {
let save_dir = match save_dir {
Some(dir) => dir,
None => get_default_save_dir().expect("Failed to get home directory"),
};
let max_retries = download_retries.unwrap_or(3);
download_model(model_id, &save_dir, max_retries).await?;
save_dir + "/" + model_id
}
};
init(common.model, model_path)?;
start_http_server(common.address, common.port).await?;
Ok(())
}
/// Run the 'serv' subcommand: start service only (no download)
async fn run_serv(args: ServArgs) -> anyhow::Result<()> {
let ServArgs {
common,
weight_path,
} = args;
init(common.model, weight_path)?;
start_http_server(common.address, common.port).await?;
Ok(())
}
/// Run the 'download' subcommand: download model only (no server)
async fn run_download(args: DownloadArgs) -> anyhow::Result<()> {
let DownloadArgs {
model,
save_dir,
download_retries,
} = args;
let model_id = get_model_id(model);
let save_dir = match save_dir {
Some(dir) => dir,
None => get_default_save_dir().expect("Failed to get home directory"),
};
let max_retries = download_retries.unwrap_or(3);
download_model(model_id, &save_dir, max_retries).await?;
Ok(())
}
/// Run the 'run' subcommand: direct model inference
fn run_run(args: RunArgs) -> anyhow::Result<()> {
use aha::exec::ExecModel;
let RunArgs {
model,
input,
output,
weight_path,
} = args;
match model {
WhichModel::MiniCPM4_0_5B => {
use aha::exec::minicpm4::MiniCPM4Exec;
MiniCPM4Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen2_5vl3B => {
use aha::exec::qwen2_5vl::Qwen2_5vlExec;
Qwen2_5vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen2_5vl7B => {
use aha::exec::qwen2_5vl::Qwen2_5vlExec;
Qwen2_5vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3_0_6B => {
use aha::exec::qwen3::Qwen3Exec;
Qwen3Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3vl2B => {
use aha::exec::qwen3vl::Qwen3vlExec;
Qwen3vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3vl4B => {
use aha::exec::qwen3vl::Qwen3vlExec;
Qwen3vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3vl8B => {
use aha::exec::qwen3vl::Qwen3vlExec;
Qwen3vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::Qwen3vl32B => {
use aha::exec::qwen3vl::Qwen3vlExec;
Qwen3vlExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::DeepSeekOCR => {
use aha::exec::deepseek_ocr::DeepSeekORExec;
DeepSeekORExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::HunyuanOCR => {
use aha::exec::hunyuan_ocr::HunyuanORExec;
HunyuanORExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::PaddleOCRVL => {
use aha::exec::paddleocr_vl::PaddleOVLExec;
PaddleOVLExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::RMBG2_0 => {
use aha::exec::rmbg2_0::RMBG2_0Exec;
RMBG2_0Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::VoxCPM => {
use aha::exec::voxcpm::VoxCPMExec;
VoxCPMExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::VoxCPM1_5 => {
use aha::exec::voxcpm1_5::VoxCPM1_5Exec;
VoxCPM1_5Exec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::GlmASRNano2512 => {
use aha::exec::glm_asr_nano::GlmASRNanoExec;
GlmASRNanoExec::run(&input, output.as_deref(), &weight_path)?;
}
WhichModel::FunASRNano2512 => {
use aha::exec::fun_asr_nano::FunASRNanoExec;
FunASRNanoExec::run(&input, output.as_deref(), &weight_path)?;
}
}
Ok(())
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
match cli.command {
Some(Commands::Cli(args)) => run_cli(args).await,
Some(Commands::Serv(args)) => run_serv(args).await,
Some(Commands::Download(args)) => run_download(args).await,
Some(Commands::Run(args)) => run_run(args),
Some(Commands::List) => run_list(),
None => {
// Backward compatibility: when no subcommand is provided, use 'cli' behavior
let model = cli.model.expect("Model is required (use -m or --model)");
let args = CliArgs {
common: CommonArgs {
address: cli.address.unwrap_or_else(|| "127.0.0.1".to_string()),
port: cli.port.unwrap_or(10100),
model,
},
weight_path: cli.weight_path,
save_dir: cli.save_dir,
download_retries: cli.download_retries,
};
run_cli(args).await
}
}
}
pub(crate) async fn start_http_server(address: String, port: u16) -> anyhow::Result<()> {
let mut builder = rocket::build().configure(Config {
address: IpAddr::from_str(&args.address)?,
port: args.port,
address: IpAddr::from_str(&address)?,
port,
limits: Limits::default()
.limit("string", ByteUnit::Mebibyte(5))
.limit("json", ByteUnit::Mebibyte(5))
@@ -128,7 +391,3 @@ pub(crate) async fn start_http_server(args: &Args) -> anyhow::Result<()> {
builder.launch().await?;
Ok(())
}
// fn main() {
// println!("Hello, world!");
// }
+15 -15
View File
@@ -34,37 +34,37 @@ use crate::models::{
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
pub enum WhichModel {
#[value(name = "minicpm4-0.5b")]
#[value(name = "minicpm4-0.5b", hide = true)]
MiniCPM4_0_5B,
#[value(name = "qwen2.5vl-3b")]
#[value(name = "qwen2.5vl-3b", hide = true)]
Qwen2_5vl3B,
#[value(name = "qwen2.5vl-7b")]
#[value(name = "qwen2.5vl-7b", hide = true)]
Qwen2_5vl7B,
#[value(name = "qwen3-0.6b")]
#[value(name = "qwen3-0.6b", hide = true)]
Qwen3_0_6B,
#[value(name = "qwen3vl-2b")]
#[value(name = "qwen3vl-2b", hide = true)]
Qwen3vl2B,
#[value(name = "qwen3vl-4b")]
#[value(name = "qwen3vl-4b", hide = true)]
Qwen3vl4B,
#[value(name = "qwen3vl-8b")]
#[value(name = "qwen3vl-8b", hide = true)]
Qwen3vl8B,
#[value(name = "qwen3vl-32b")]
#[value(name = "qwen3vl-32b", hide = true)]
Qwen3vl32B,
#[value(name = "deepseek-ocr")]
#[value(name = "deepseek-ocr", hide = true)]
DeepSeekOCR,
#[value(name = "hunyuan-ocr")]
#[value(name = "hunyuan-ocr", hide = true)]
HunyuanOCR,
#[value(name = "paddleocr-vl")]
#[value(name = "paddleocr-vl", hide = true)]
PaddleOCRVL,
#[value(name = "rmbg2.0")]
RMBG2_0,
#[value(name = "voxcpm")]
#[value(name = "voxcpm", hide = true)]
VoxCPM,
#[value(name = "voxcpm1.5")]
#[value(name = "voxcpm1.5", hide = true)]
VoxCPM1_5,
#[value(name = "glm-asr-nano-2512")]
#[value(name = "glm-asr-nano-2512", hide = true)]
GlmASRNano2512,
#[value(name = "fun-asr-nano-2512")]
#[value(name = "fun-asr-nano-2512", hide = true)]
FunASRNano2512,
}
+92
View File
@@ -0,0 +1,92 @@
use serde::Deserialize;
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct FunASRNanoConfig {
pub audio_encoder_conf: AudioEncoderConf,
pub llm_conf: LlmConf,
pub audio_adaptor_conf: AudioAdaptorConf,
pub detach_ctc_decoder: bool,
pub ctc_decoder_conf: CtcDecoderConf,
pub ctc_weight: f64,
pub ctc_conf: CtcConf,
pub frontend_conf: FrontendConf,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct AudioEncoderConf {
pub output_size: usize,
pub attention_heads: usize,
pub linear_units: usize,
pub num_blocks: usize,
pub tp_blocks: usize,
pub dropout_rate: f64,
pub positional_dropout_rate: f64,
pub attention_dropout_rate: f64,
pub input_layer: String,
pub pos_enc_class: String,
pub normalize_before: bool,
pub kernel_size: usize,
pub sanm_shfit: usize,
pub selfattention_layer_type: String,
pub freeze: bool,
pub freeze_layer_num: i32,
pub feat_permute: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct LlmConf {
pub hub: String,
pub freeze: bool,
pub llm_dtype: String,
pub init_param_path: String,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct AudioAdaptorConf {
pub downsample_rate: usize,
pub use_low_frame_rate: bool,
pub ffn_dim: usize,
pub llm_dim: usize,
pub encoder_dim: usize,
pub n_layer: usize,
pub freeze: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct CtcDecoderConf {
pub downsample_rate: u32,
pub ffn_dim: u32,
pub llm_dim: u32,
pub encoder_dim: u32,
pub n_layer: u32,
pub freeze: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct CtcConf {
pub dropout_rate: f64,
pub ctc_type: String,
pub reduce: bool,
pub ignore_nan_grad: bool,
}
#[derive(Debug, Clone, PartialEq, Deserialize)]
pub struct FrontendConf {
pub fs: usize,
pub window: String,
pub n_mels: usize,
pub frame_length: f32,
pub frame_shift: f32,
pub lfr_m: usize,
pub lfr_n: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub cmvn_file: Option<serde_yaml::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct Qwen3ASRGenerationConfig {
pub do_sample: bool,
pub eos_token_id: Vec<usize>,
pub pad_token_id: usize,
pub temperature: f32,
}
+212
View File
@@ -0,0 +1,212 @@
use std::collections::HashMap;
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use rocket::async_stream::stream;
use rocket::futures::Stream;
use crate::{
chat_template::ChatTemplate, models::{
GenerateModel,
fun_asr_nano::{
config::FunASRNanoConfig, model::FunAsrNanoModel,
},
qwen3::config::{Qwen3Config, Qwen3GenerationConfig}, qwen3_asr::{config::Qwen3ASRGenerationConfig, processor::Qwen3AsrProcessor},
}, tokenizer::TokenizerModel, utils::{
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
get_dtype, get_logit_processor,
}
};
pub struct Qwen3AsrGenerateModel<'a> {
chat_template: ChatTemplate<'a>,
// tokenizer: TokenizerModel,
processor: Qwen3AsrProcessor,
// fun_asr_nano: FunAsrNanoModel,
device: Device,
// dtype: DType,
eos_token_id1: u32,
eos_token_id2: u32,
generation_config: Qwen3ASRGenerationConfig,
model_name: String,
}
impl<'a> Qwen3AsrGenerateModel<'a> {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let chat_template = ChatTemplate::init(path)?;
let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3ASRGenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?;
let device = get_device(device);
let processor = Qwen3AsrProcessor::new(&device)?;
Ok(Self {
chat_template,
// tokenizer,
processor,
// fun_asr_nano,
device,
// dtype,
eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config,
model_name: "qwen3-asr".to_string(),
})
}
pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> {
let temperature = match mes.temperature {
None => self.generation_config.temperature,
Some(tem) => tem,
};
let seed = match mes.seed {
None => 34562u64,
Some(s) => s as u64,
};
let mut logit_processor =
get_logit_processor(Some(temperature), mes.top_p, None, seed);
let render_text = self.chat_template.apply_chat_template(&mes)?;
let audio_data =
self.processor.process_info(&mes, &render_text)?;
// for audio in audio_data {
// let text =
// }
Ok(())
}
}
// impl<'a> GenerateModel for Qwen3AsrGenerateModel<'a> {
// fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
// let temperature = match mes.temperature {
// None => self.generation_config.temperature,
// Some(tem) => tem,
// };
// let seed = match mes.seed {
// None => 34562u64,
// Some(s) => s as u64,
// };
// let mut logit_processor =
// get_logit_processor(Some(temperature), mes.top_p, None, seed);
// let audio_data =
// self.processor.process_info(&mes)?;
// for audio in audio_data {
// let text =
// }
// let mut speech = Some(speech.to_dtype(self.dtype)?);
// let mut fbank_mask = Some(&fbank_mask);
// let mut seq_len = input_ids.dim(1)?;
// let mut seqlen_offset = 0;
// let mut generate = Vec::new();
// let sample_len = mes.max_tokens.unwrap_or(1024);
// for _ in 0..sample_len {
// let logits = self.fun_asr_nano.forward(
// &input_ids,
// speech.as_ref(),
// fbank_mask,
// 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_id1 || next_token == self.eos_token_id2 {
// break;
// }
// seqlen_offset += seq_len;
// seq_len = 1;
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
// speech = None;
// fbank_mask = None;
// }
// let num_token = generate.len() as u32;
// let res = self.tokenizer.token_decode(generate)?;
// self.fun_asr_nano.clear_kv_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 temperature = match mes.temperature {
// None => self.generation_config.temperature,
// Some(tem) => tem,
// };
// let top_p = match mes.top_p {
// None => self.generation_config.top_p,
// Some(top_p) => top_p,
// };
// let top_k = self.generation_config.top_k;
// let seed = match mes.seed {
// None => 34562u64,
// Some(s) => s as u64,
// };
// let mut logit_processor =
// get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
// let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
// 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 speech = Some(speech.to_dtype(self.dtype)?);
// let mut fbank_mask = Some(&fbank_mask);
// let mut input_ids = input_ids;
// for _ in 0..sample_len {
// let logits = self.fun_asr_nano.forward(
// &input_ids,
// speech.as_ref(),
// fbank_mask,
// 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)?;
// speech = None;
// fbank_mask = None;
// continue;
// }
// error_tokens.clear();
// let chunk = build_completion_chunk_response(decoded_token, &self.model_name, None, None);
// yield Ok(chunk);
// if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
// break;
// }
// seqlen_offset += seq_len;
// seq_len = 1;
// input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
// speech = None;
// fbank_mask = None;
// }
// self.fun_asr_nano.clear_kv_cache();
// };
// Ok(Box::new(Box::pin(stream)))
// }
// }
+4
View File
@@ -0,0 +1,4 @@
pub mod config;
pub mod generate;
pub mod model;
pub mod processor;
+646
View File
@@ -0,0 +1,646 @@
use anyhow::Result;
use candle_core::{D, IndexOp, Tensor};
use candle_nn::{Conv1d, LayerNorm, Linear, Module, VarBuilder, linear, ops::softmax_last_dim};
use crate::{
models::{
common::{
NaiveAttention, TwoLinearMLP, eager_attention_forward, get_conv1d, get_layer_norm,
},
fun_asr_nano::config::FunASRNanoConfig,
qwen3::{config::Qwen3Config, model::Qwen3Model},
},
position_embed::sinusoidal_pe::SinusoidalPositionEncoderCat,
utils::tensor_utils::{get_equal_mask, mask_filled, masked_scatter_dim0},
};
pub struct MultiHeadedAttentionSANM {
head_dim: usize,
n_head: usize,
linear_out: Linear,
linear_q_k_v: Linear,
fsmn_block: Conv1d,
left_padding: usize,
right_padding: usize,
scaling: f64,
}
impl MultiHeadedAttentionSANM {
pub fn new(
vb: VarBuilder,
n_head: usize,
in_dim: usize,
hidden_dim: usize,
kernel_size: usize,
sanm_shfit: usize,
) -> Result<Self> {
let head_dim = hidden_dim / n_head;
let linear_out = linear(hidden_dim, hidden_dim, vb.pp("linear_out"))?;
let linear_q_k_v = linear(in_dim, hidden_dim * 3, vb.pp("linear_q_k_v"))?;
let fsmn_block = get_conv1d(
vb.pp("fsmn_block"),
hidden_dim,
hidden_dim,
kernel_size,
0,
1,
1,
hidden_dim,
false,
)?;
let mut left_padding = (kernel_size - 1) / 2;
if sanm_shfit > 0 {
left_padding += sanm_shfit;
}
let right_padding = kernel_size - 1 - left_padding;
let scaling = (head_dim as f64).powf(-0.5);
Ok(Self {
head_dim,
n_head,
linear_out,
linear_q_k_v,
fsmn_block,
left_padding,
right_padding,
scaling,
})
}
pub fn forward_fsmn(
&self,
inputs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
) -> Result<Tensor> {
let mut inputs = inputs.clone();
let mask = if let Some(mask) = mask {
let mut mask = mask.unsqueeze(D::Minus1)?.unsqueeze(0)?;
if let Some(mask_shfit_chunk) = mask_shfit_chunk {
mask = mask.broadcast_mul(mask_shfit_chunk)?;
}
inputs = inputs.broadcast_mul(&mask)?;
Some(mask)
} else {
None
};
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 = xs.transpose(1, 2)?;
let mut xs = xs.add(&inputs)?;
if let Some(mask) = mask {
xs = xs.broadcast_mul(&mask)?;
}
Ok(xs)
}
pub fn forward_qkv(&self, xs: &Tensor) -> Result<(Tensor, Tensor, Tensor, Tensor)> {
let (b, t, _) = xs.dims3()?;
let q_k_v = self
.linear_q_k_v
.forward(xs)?
.reshape((b, t, 3, self.n_head, ()))?
.permute((2, 0, 3, 1, 4))?
.contiguous()?;
let q_h = q_k_v.i(0)?.contiguous()?;
let k_h = q_k_v.i(1)?.contiguous()?;
let v_h = q_k_v.i(2)?.contiguous()?;
let v = v_h.transpose(1, 2)?.reshape((b, t, ()))?;
Ok((q_h, k_h, v_h, v))
}
pub fn forward_attention(
&self,
values: &Tensor,
scores: &Tensor,
mask: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let bs = scores.dim(0)?;
let attn = if let Some(mask) = mask {
let mask = if let Some(mask_att_chunk_encoder) = mask_att_chunk_encoder {
mask.mul(mask_att_chunk_encoder)?
} else {
mask.clone()
};
// mask: rank = 2
let mask = get_equal_mask(&mask, 0)?;
let scores = mask_filled(scores, &mask, f32::NEG_INFINITY)?;
let attn = softmax_last_dim(&scores)?;
mask_filled(&attn, &mask, 0.0)?
} else {
softmax_last_dim(scores)?
};
let xs = attn.matmul(values)?;
let xs =
xs.transpose(1, 2)?
.contiguous()?
.reshape((bs, (), self.n_head * self.head_dim))?;
let xs = self.linear_out.forward(&xs)?;
Ok(xs)
}
pub fn forward_simple(&self, xs: &Tensor) -> Result<Tensor> {
let (b, t, _) = xs.dims3()?;
let q_k_v = self.linear_q_k_v.forward(xs)?;
let dim = self.head_dim * self.n_head;
let q_h = q_k_v
.narrow(D::Minus1, 0, dim)?
.reshape((b, t, self.n_head, ()))?
.permute((0, 2, 1, 3))?;
let k_h = q_k_v
.narrow(D::Minus1, dim, dim)?
.reshape((b, t, self.n_head, ()))?
.permute((0, 2, 1, 3))?;
let v = q_k_v.narrow(D::Minus1, dim * 2, dim)?;
let v_h = v.reshape((b, t, self.n_head, ()))?.permute((0, 2, 1, 3))?;
let fsmn_memory = v.transpose(1, 2)?;
let fsmn_memory = fsmn_memory
.pad_with_zeros(D::Minus1, self.left_padding, self.right_padding)?
.contiguous()?;
let fsmn_memory = self.fsmn_block.forward(&fsmn_memory)?;
// let fsmn_memory = conv1d_group_parallel(&fsmn_memory, &self.fsmn_block)?;
let fsmn_memory = fsmn_memory.transpose(1, 2)?;
let fsmn_memory = fsmn_memory.add(&v)?;
let att_outs = eager_attention_forward(&q_h, &k_h, &v_h, None, None, self.scaling)?;
let att_outs = att_outs.reshape((b, t, ()))?;
let att_outs = self.linear_out.forward(&att_outs)?;
let att_outs = att_outs.add(&fsmn_memory)?;
Ok(att_outs)
}
pub fn forward(
&self,
xs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let (q_h, k_h, v_h, v) = self.forward_qkv(xs)?;
let fsmn_memory = self.forward_fsmn(&v, mask, mask_shfit_chunk)?;
let q_h = q_h.affine(self.scaling, 0.0)?;
let scores = q_h.matmul(&k_h.transpose(D::Minus2, D::Minus1)?)?;
let attn_outs = self.forward_attention(&v_h, &scores, mask, mask_att_chunk_encoder)?;
let att_outs = attn_outs.add(&fsmn_memory)?;
Ok(att_outs)
}
}
pub struct EncoderLayerSANM {
self_attn: MultiHeadedAttentionSANM,
feed_forward: TwoLinearMLP,
norm1: LayerNorm,
norm2: LayerNorm,
concat_linear: Option<Linear>,
normalize_before: bool,
in_dim: usize,
hidden_dim: usize,
}
impl EncoderLayerSANM {
pub fn new(
vb: VarBuilder,
in_dim: usize,
hidden_dim: usize,
n_head: usize,
kernel_size: usize,
sanm_shfit: usize,
hidden_units: usize,
normalize_before: bool,
concat_after: bool,
) -> Result<Self> {
let self_attn = MultiHeadedAttentionSANM::new(
vb.pp("self_attn"),
n_head,
in_dim,
hidden_dim,
kernel_size,
sanm_shfit,
)?;
let feed_forward = TwoLinearMLP::new(
vb.pp("feed_forward"),
hidden_dim,
hidden_units,
hidden_dim,
candle_nn::Activation::Relu,
true,
"w_1",
"w_2",
)?;
let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, in_dim)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, hidden_dim)?;
let concat_linear = if concat_after {
let lin = linear(hidden_dim * 2, hidden_dim, vb.pp("concat_linear"))?;
Some(lin)
} else {
None
};
Ok(Self {
self_attn,
feed_forward,
norm1,
norm2,
concat_linear,
normalize_before,
in_dim,
hidden_dim,
})
}
pub fn forward(
&self,
xs: &Tensor,
mask: Option<&Tensor>,
mask_shfit_chunk: Option<&Tensor>,
mask_att_chunk_encoder: Option<&Tensor>,
) -> Result<Tensor> {
let stoch_layer_coeff = 1.0f64;
let residual = xs.clone();
let mut xs = if self.normalize_before {
self.norm1.forward(xs)?
} else {
xs.clone()
};
if self.concat_linear.is_some() {
let attn =
self.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
if self.in_dim == self.hidden_dim {
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
} else {
xs = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
}
} else if self.in_dim == self.hidden_dim {
let attn = self
.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&attn)?;
} else {
xs = self
.self_attn
.forward(&xs, mask, mask_shfit_chunk, mask_att_chunk_encoder)?
.affine(stoch_layer_coeff, 0.0)?;
}
if !self.normalize_before {
xs = self.norm1.forward(&xs)?;
}
let residual = xs.clone();
if self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
xs = self
.feed_forward
.forward(&xs)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&xs)?;
if !self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
Ok(xs)
}
pub fn forward_simple(&self, xs: &Tensor) -> Result<Tensor> {
let residual = xs.clone();
let mut xs = self.norm1.forward(xs)?;
if self.in_dim == self.hidden_dim {
let attn = self.self_attn.forward_simple(&xs)?;
xs = residual.add(&attn)?;
} else {
xs = self.self_attn.forward_simple(&xs)?;
}
let residual = xs.clone();
let xs = self.norm2.forward(&xs)?;
let xs = self.feed_forward.forward(&xs)?;
let xs = residual.add(&xs)?;
Ok(xs)
}
}
pub struct SenseVoiceEncoderSmall {
embed: SinusoidalPositionEncoderCat,
encoders0: EncoderLayerSANM,
encoders: Vec<EncoderLayerSANM>,
tp_encoders: Vec<EncoderLayerSANM>,
after_norm: LayerNorm,
tp_norm: LayerNorm,
scaling: f64,
}
impl SenseVoiceEncoderSmall {
pub fn new(
vb: VarBuilder,
input_size: usize,
output_size: usize,
attention_heads: usize,
linear_units: usize,
num_blocks: usize,
tp_blocks: usize,
normalize_before: bool,
kernel_size: usize,
sanm_shfit: usize,
) -> Result<Self> {
let embed = SinusoidalPositionEncoderCat::new(Some(input_size), true, vb.device())?;
let encoders0 = EncoderLayerSANM::new(
vb.pp("encoders0.0"),
input_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
let mut encoders = vec![];
let vb_encoders = vb.pp("encoders");
for i in 0..(num_blocks - 1) {
let encoder_i = EncoderLayerSANM::new(
vb_encoders.pp(i),
output_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
encoders.push(encoder_i);
}
let vb_tp_encoders = vb.pp("tp_encoders");
let mut tp_encoders = vec![];
for i in 0..tp_blocks {
let tp_blocks_i = EncoderLayerSANM::new(
vb_tp_encoders.pp(i),
output_size,
output_size,
attention_heads,
kernel_size,
sanm_shfit,
linear_units,
normalize_before,
false,
)?;
tp_encoders.push(tp_blocks_i);
}
let after_norm = get_layer_norm(vb.pp("after_norm"), 1e-5, output_size)?;
let tp_norm = get_layer_norm(vb.pp("tp_norm"), 1e-5, output_size)?;
let scaling = (output_size as f64).powf(0.5);
Ok(Self {
embed,
encoders0,
encoders,
tp_encoders,
after_norm,
tp_norm,
scaling,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let xs = xs.affine(self.scaling, 0.0)?;
let xs = self.embed.forward(&xs, 0)?;
let mut xs = self.encoders0.forward_simple(&xs)?;
for encoder_layer in &self.encoders {
xs = encoder_layer.forward_simple(&xs)?;
}
xs = self.after_norm.forward(&xs)?;
for tp_layer in &self.tp_encoders {
xs = tp_layer.forward_simple(&xs)?;
}
xs = self.tp_norm.forward(&xs)?;
Ok(xs)
}
}
pub struct AdaptorEncoderLayer {
self_attn: NaiveAttention,
feed_forward: TwoLinearMLP,
norm1: LayerNorm,
norm2: LayerNorm,
concat_linear: Option<Linear>,
normalize_before: bool,
}
impl AdaptorEncoderLayer {
pub fn new(
vb: VarBuilder,
llm_dim: usize,
n_head: usize,
normalize_before: bool,
concat_after: bool,
) -> Result<Self> {
let self_attn = NaiveAttention::new(
vb.pp("self_attn"),
llm_dim,
n_head,
n_head,
None,
true,
Some("linear_q"),
Some("linear_k"),
Some("linear_v"),
Some("linear_out"),
)?;
let feed_forward = TwoLinearMLP::new(
vb.pp("feed_forward"),
llm_dim,
llm_dim / 4,
llm_dim,
candle_nn::Activation::Relu,
true,
"w_1",
"w_2",
)?;
let norm1 = get_layer_norm(vb.pp("norm1"), 1e-5, llm_dim)?;
let norm2 = get_layer_norm(vb.pp("norm2"), 1e-5, llm_dim)?;
let concat_linear = if concat_after {
let lin = linear(llm_dim * 2, llm_dim, vb.pp("concat_linear"))?;
Some(lin)
} else {
None
};
Ok(Self {
self_attn,
feed_forward,
norm1,
norm2,
concat_linear,
normalize_before,
})
}
pub fn forward(&self, xs: &Tensor, mask: Option<&Tensor>) -> Result<Tensor> {
let stoch_layer_coeff = 1.0f64;
let residual = xs.clone();
let mut xs = if self.normalize_before {
self.norm1.forward(xs)?
} else {
xs.clone()
};
if self.concat_linear.is_some() {
let attn = self.self_attn.forward(&xs, None, None, mask, false)?;
let x_concat = Tensor::cat(&[&xs, &attn], D::Minus1)?;
let x_concat = self
.concat_linear
.as_ref()
.unwrap()
.forward(&x_concat)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&x_concat)?;
} else {
let attn = self
.self_attn
.forward(&xs, None, None, mask, false)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&attn)?;
}
if !self.normalize_before {
xs = self.norm1.forward(&xs)?;
}
let residual = xs.clone();
if self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
xs = self
.feed_forward
.forward(&xs)?
.affine(stoch_layer_coeff, 0.0)?;
xs = residual.add(&xs)?;
if !self.normalize_before {
xs = self.norm2.forward(&xs)?;
}
Ok(xs)
}
}
pub struct AudioAdaptor {
k: usize,
linear1: Linear,
linear2: Linear,
blocks: Vec<AdaptorEncoderLayer>,
}
impl AudioAdaptor {
pub fn new(
vb: VarBuilder,
downsample_rate: usize,
encoder_dim: usize,
llm_dim: usize,
ffn_dim: usize,
n_layer: usize,
attention_heads: usize,
) -> Result<Self> {
let linear1 = linear(encoder_dim * downsample_rate, ffn_dim, vb.pp("linear1"))?;
let linear2 = linear(ffn_dim, llm_dim, vb.pp("linear2"))?;
let mut blocks = vec![];
let vb_blocks = vb.pp("blocks");
for i in 0..n_layer {
let layer =
AdaptorEncoderLayer::new(vb_blocks.pp(i), llm_dim, attention_heads, true, false)?;
blocks.push(layer);
}
Ok(Self {
k: downsample_rate,
linear1,
linear2,
blocks,
})
}
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let (bs, seq_len, dim) = xs.dims3()?;
let chunk_num = (seq_len - 1) / self.k + 1;
let pad_num = chunk_num * self.k - seq_len;
let xs = xs.pad_with_zeros(1, 0, pad_num)?;
let xs = xs.contiguous()?.reshape((bs, chunk_num, dim * self.k))?;
let xs = self.linear1.forward(&xs)?.relu()?;
let mut xs = self.linear2.forward(&xs)?;
for block in &self.blocks {
xs = block.forward(&xs, None)?;
}
Ok(xs)
}
}
pub struct FunAsrNanoModel {
audio_encoder: SenseVoiceEncoderSmall,
audio_adaptor: AudioAdaptor,
llm: Qwen3Model,
}
impl FunAsrNanoModel {
pub fn new(vb: VarBuilder, config: &FunASRNanoConfig, llm_cfg: &Qwen3Config) -> Result<Self> {
let input_size = config.frontend_conf.lfr_m * config.frontend_conf.n_mels;
let audio_encoder = SenseVoiceEncoderSmall::new(
vb.pp("audio_encoder"),
input_size,
config.audio_encoder_conf.output_size,
config.audio_encoder_conf.attention_heads,
config.audio_encoder_conf.linear_units,
config.audio_encoder_conf.num_blocks,
config.audio_encoder_conf.tp_blocks,
config.audio_encoder_conf.normalize_before,
config.audio_encoder_conf.kernel_size,
config.audio_encoder_conf.sanm_shfit,
)?;
let audio_adaptor = AudioAdaptor::new(
vb.pp("audio_adaptor"),
config.audio_adaptor_conf.downsample_rate,
config.audio_adaptor_conf.encoder_dim,
config.audio_adaptor_conf.llm_dim,
config.audio_adaptor_conf.ffn_dim,
config.audio_adaptor_conf.n_layer,
8,
)?;
let llm = Qwen3Model::new(llm_cfg, vb.pp("llm"))?;
Ok(Self {
audio_encoder,
audio_adaptor,
llm,
})
}
pub fn forward(
&mut self,
input_ids: &Tensor,
speech: Option<&Tensor>,
fbank_mask: Option<&Tensor>,
seqlen_offset: usize,
) -> Result<Tensor> {
let mut inputs_embeds = self.llm.embedding_token_id(input_ids)?;
if let Some(speech) = speech
&& let Some(fbank_mask) = fbank_mask
{
let speech = self.audio_encoder.forward(speech)?;
let encoder_out = self.audio_adaptor.forward(&speech)?;
let speech_token_len = fbank_mask.sum_all()?.to_scalar::<u32>()?;
let audio_embed = encoder_out
.squeeze(0)?
.narrow(0, 0, speech_token_len as usize)?;
inputs_embeds = masked_scatter_dim0(&inputs_embeds, &audio_embed, fbank_mask)?;
}
let logits = self
.llm
.forward(None, Some(&inputs_embeds), seqlen_offset)?;
Ok(logits)
}
pub fn clear_kv_cache(&mut self) {
self.llm.clear_kv_cache();
}
}
+127
View File
@@ -0,0 +1,127 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use candle_core::{Device, Tensor};
use serde_json::{Value, json};
use crate::utils::{
audio_utils::{extract_audios, split_audio_into_chunks},
capitalize_first_letter, extract_user_text_vec,
tensor_utils::float_range_normalize,
};
pub struct Qwen3AsrProcessor {
device: Device,
sample_rate: usize,
support_language: Vec<String>,
max_asr_input_seconds: f32,
}
impl Qwen3AsrProcessor {
pub fn new(device: &Device) -> Result<Self> {
let support_language: Vec<String> = vec![
"Chinese",
"English",
"Cantonese",
"Arabic",
"German",
"French",
"Spanish",
"Portuguese",
"Indonesian",
"Italian",
"Korean",
"Russian",
"Thai",
"Vietnamese",
"Japanese",
"Turkish",
"Hindi",
"Malay",
"Dutch",
"Swedish",
"Danish",
"Finnish",
"Polish",
"Czech",
"Filipino",
"Persian",
"Greek",
"Romanian",
"Hungarian",
"Macedonian",
]
.iter()
.map(|s| s.to_string())
.collect();
Ok(Self {
device: device.clone(),
sample_rate: 16000,
support_language,
max_asr_input_seconds: 1200.0,
})
}
pub fn process_audio(&self, mes: &ChatCompletionParameters) -> Result<Vec<Tensor>> {
let audio_tensors = extract_audios(mes, &self.device, Some(self.sample_rate))?;
audio_tensors
.iter()
.map(|audio| float_range_normalize(&audio))
.collect()
}
pub fn validate_language(&self, lang: &String) -> bool {
self.support_language.contains(lang)
}
pub fn process_info(&self, mes: &ChatCompletionParameters, render: &str) -> Result<()> {
let audio_count = render
.matches("<|audio_start|><|audio_pad|><|audio_end|>")
.count();
let mut render = if audio_count > 1 {
render.replace(
&"<|audio_start|><|audio_pad|><|audio_end|>".repeat(audio_count),
"<|audio_start|><|audio_pad|><|audio_end|>",
)
} else {
render.to_string()
};
if let Some(map) = &mes.metadata
&& map.contains_key("language")
{
let lang = map.get("language").unwrap();
let lang = capitalize_first_letter(lang);
if self.validate_language(&lang) {
render = format!("{}language {}'<asr_text>'", render, lang);
}
}
let audio_tensors = self.process_audio(mes)?;
let audio_len = audio_tensors.len();
if audio_len != audio_count {
return Err(anyhow::anyhow!("audio_pad num != audio num"));
}
let mut split_wavs = vec![];
for wav in audio_tensors.iter() {
let wavs = split_audio_into_chunks(wav, self.sample_rate, self.max_asr_input_seconds)?;
split_wavs.extend_from_slice(&wavs);
}
// let mut audio_datas = vec![];
// for (i, wav) in audio_tensors.iter().enumerate() {
// let wavs = split_audio_into_chunks(wav, self.sample_rate, self.max_asr_input_seconds)?;
// for i_w in wavs {
// let audio_data = AudioData {
// wav: i_w,
// language: langs[i].clone(),
// };
// audio_datas.push(audio_data);
// }
// }
Ok(())
}
}
pub struct AudioData {
pub wav: Tensor,
pub language: Option<String>,
}
+15 -1
View File
@@ -4,7 +4,7 @@ pub mod tensor_utils;
pub mod video_utils;
use std::io::Read;
use std::{collections::HashMap, fs, process::Command, time::Duration};
use std::{collections::HashMap, fs, path::PathBuf, process::Command, time::Duration};
use aha_openai_dive::v1::resources::{
chat::{
@@ -758,3 +758,17 @@ pub async fn download_model(
}
}
}
pub fn get_file_path(file: &str) -> Result<PathBuf> {
let path = url::Url::parse(file)?;
let path = path.to_file_path();
let path = match path {
Ok(path) => path,
Err(_) => {
let mut path = file.to_owned();
path = path.split_off(7);
PathBuf::from(path)
}
};
Ok(path)
}
+94
View File
@@ -0,0 +1,94 @@
use std::{pin::pin, time::Instant};
use aha::models::{GenerateModel, fun_asr_nano::generate::FunAsrNanoGenerateModel, qwen3_asr::generate::Qwen3AsrGenerateModel};
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use rocket::futures::StreamExt;
#[test]
fn qwen3_asr_generate() -> Result<()> {
// RUST_BACKTRACE=1 cargo test -F cuda qwen3_asr_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-ASR-0.6B/", save_dir);
let message = r#"
{
"model": "qwen3-asr",
"messages": [
{
"role": "user",
"content": [
{
"type": "audio",
"audio_url":
{
"url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav"
}
}
]
}
],
"metadata": {"language": "Chinese"}
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
let i_start = Instant::now();
let mut model = Qwen3AsrGenerateModel::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 = model.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_asr_stream() -> Result<()> {
// // RUST_BACKTRACE=1 cargo test -F cuda fun_asr_nano_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-ASR-0.6B/", save_dir);
// let message = r#"
// {
// "model": "qwen3-asr",
// "messages": [
// {
// "role": "user",
// "content": [
// {
// "type": "audio",
// "audio_url":
// {
// "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav"
// }
// },
// {
// "type": "text",
// "text": "语音转写:"
// }
// ]
// }
// ]
// }
// "#;
// let mes: ChatCompletionParameters = serde_json::from_str(message)?;
// let i_start = Instant::now();
// let mut fun_asr_model = FunAsrNanoGenerateModel::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!(fun_asr_model.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(())
// }
+52
View File
@@ -5,6 +5,58 @@ use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::Result;
use rocket::futures::StreamExt;
#[test]
fn qwen3vl_thinking_generate() -> Result<()> {
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3vl_thinking_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-VL-2B-Thinking/", save_dir);
let message = r#"
{
"model": "qwen3vl-thinking",
"messages": [
{
"role": "user",
"content": [
{
"type": "image",
"image_url":
{
"url": "file://./assets/img/ocr_test1.png"
}
},
{
"type": "text",
"text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本"
}
]
}
],
"max_tokens": 10240
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
let i_start = Instant::now();
let mut qwen3vl = Qwen3VLGenerateModel::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(())
}
#[test]
fn qwen3vl_generate() -> Result<()> {
// test with cuda: RUST_BACKTRACE=1 cargo test -F cuda,ffmpeg qwen3vl_generate -r -- --nocapture