diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..196cfc7 --- /dev/null +++ b/CHANGELOG.md @@ -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 ` - Specify which model to use + - `-in, --input ` - Input text or file path (model-specific interpretation) + - `-out, --output ` - Output file path (optional, auto-generated if not specified) + - `--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
` - Server address (default: 127.0.0.1) +- `-p, --port ` - Server port (default: 10100) +- `-m, --model ` - Model to use (required) +- `--weight-path ` - Local model weight path (optional) +- `--save-dir ` - Directory to save downloaded model (optional) +- `--download-retries ` - Download retry attempts (default: 3) + +`aha serv`: +- `-a, --address
` - Server address (default: 127.0.0.1) +- `-p, --port ` - Server port (default: 10100) +- `-m, --model ` - Model to use (required) +- `--weight-path ` - Local model weight path (required) + +`aha download`: +- `-m, --model ` - Model to download (required) +- `-s, --save-dir ` - Directory to save downloaded model (optional) +- `--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 diff --git a/Makefile b/Makefile index 6369fb8..a69a21e 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/docs/CLI_USAGE.md b/docs/CLI_USAGE.md new file mode 100644 index 0000000..7a5a0db --- /dev/null +++ b/docs/CLI_USAGE.md @@ -0,0 +1,291 @@ +# AHA 命令行使用说明 + +## 概述 + +AHA 是一个基于 Candle 框架的高性能模型推理库,支持多种多模态模型,包括视觉、语言和语音模型。 + +```bash +aha [COMMAND] [OPTIONS] +``` + +## 全局选项 + +| 选项 | 说明 | 默认值 | +|------|------|--------| +| `-a, --address
` | 服务监听地址 | 127.0.0.1 | +| `-p, --port ` | 服务监听端口 | 10100 | +| `-m, --model ` | 模型类型(必选) | - | +| `--weight-path ` | 本地模型权重路径 | - | +| `--save-dir ` | 模型下载保存目录 | ~/.aha/ | +| `--download-retries ` | 下载重试次数 | 3 | +| `-h, --help` | 显示帮助信息 | - | +| `-V, --version` | 显示版本号 | - | + +## 子命令 + +### cli - 下载模型并启动服务(默认) + +下载指定的模型并启动 HTTP 服务。当不指定子命令时,默认使用此命令。 + +**语法:** +```bash +aha cli [OPTIONS] --model +``` + +**选项:** + +| 选项 | 说明 | 默认值 | +|------|------|--------| +| `-a, --address
` | 服务监听地址 | 127.0.0.1 | +| `-p, --port ` | 服务监听端口 | 10100 | +| `-m, --model ` | 模型类型(必选) | - | +| `--weight-path ` | 本地模型权重路径(如指定则跳过下载) | - | +| `--save-dir ` | 模型下载保存目录 | ~/.aha/ | +| `--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 --input [--input ] --weight-path +``` + +**选项:** + +| 选项 | 说明 | 默认值 | +|------|------|--------| +| `-m, --model ` | 模型类型(必选) | - | +| `-i, --input ` | 输入文本或文件路径(模型特定解释,支持1-2个参数, input1: 提示文本, input2: 文件地址) | - | +| `-o, --output ` | 输出文件路径(可选,未指定则自动生成) | - | +| `--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 --weight-path +``` + +**选项:** + +| 选项 | 说明 | 默认值 | +|------|------|--------| +| `-a, --address
` | 服务监听地址 | 127.0.0.1 | +| `-p, --port ` | 服务监听端口 | 10100 | +| `-m, --model ` | 模型类型(必选) | - | +| `--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 +``` + +**选项:** + +| 选项 | 说明 | 默认值 | +|------|------|--------| +| `-m, --model ` | 模型类型(必选) | - | +| `-s, --save-dir ` | 模型下载保存目录 | ~/.aha/ | +| `--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 +``` diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs index 48aefb4..0cbf6d1 100644 --- a/src/chat_template/mod.rs +++ b/src/chat_template/mod.rs @@ -19,12 +19,12 @@ pub fn get_template(path: String) -> Result { // 修复模板中的问题行 let fixed_template = chat_template .replace( - "message.content.startswith('')", - "message.content is startingwith('')", // 使用minijinja中的 is startingwith 替换 + "content.startswith('')", + "content is startingwith('')", // 使用minijinja中的 is startingwith 替换 ) .replace( - "message.content.endswith('')", - "message.content is endingwith('')", // 使用minijinja中的 is endingwith 替换 + "content.endswith('')", + "content is endingwith('')", // 使用minijinja中的 is endingwith 替换 ) .replace( "content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')", diff --git a/src/exec/deepseek_ocr.rs b/src/exec/deepseek_ocr.rs new file mode 100644 index 0000000..ae46595 --- /dev/null +++ b/src/exec/deepseek_ocr.rs @@ -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": "\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(()) + } +} diff --git a/src/exec/fun_asr_nano.rs b/src/exec/fun_asr_nano.rs new file mode 100644 index 0000000..1af00fe --- /dev/null +++ b/src/exec/fun_asr_nano.rs @@ -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(()) + } +} diff --git a/src/exec/glm_asr_nano.rs b/src/exec/glm_asr_nano.rs new file mode 100644 index 0000000..70c1d3a --- /dev/null +++ b/src/exec/glm_asr_nano.rs @@ -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(()) + } +} diff --git a/src/exec/hunyuan_ocr.rs b/src/exec/hunyuan_ocr.rs new file mode 100644 index 0000000..3eb5a7b --- /dev/null +++ b/src/exec/hunyuan_ocr.rs @@ -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(()) + } +} diff --git a/src/exec/minicpm4.rs b/src/exec/minicpm4.rs new file mode 100644 index 0000000..5e36a4f --- /dev/null +++ b/src/exec/minicpm4.rs @@ -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(()) + } +} diff --git a/src/exec/mod.rs b/src/exec/mod.rs new file mode 100644 index 0000000..67d9735 --- /dev/null +++ b/src/exec/mod.rs @@ -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<()>; +} diff --git a/src/exec/paddleocr_vl.rs b/src/exec/paddleocr_vl.rs new file mode 100644 index 0000000..803f0eb --- /dev/null +++ b/src/exec/paddleocr_vl.rs @@ -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(()) + } +} diff --git a/src/exec/qwen2_5vl.rs b/src/exec/qwen2_5vl.rs new file mode 100644 index 0000000..ecaaa2e --- /dev/null +++ b/src/exec/qwen2_5vl.rs @@ -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(()) + } +} diff --git a/src/exec/qwen3.rs b/src/exec/qwen3.rs new file mode 100644 index 0000000..7a7d460 --- /dev/null +++ b/src/exec/qwen3.rs @@ -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(()) + } +} diff --git a/src/exec/qwen3vl.rs b/src/exec/qwen3vl.rs new file mode 100644 index 0000000..a336d24 --- /dev/null +++ b/src/exec/qwen3vl.rs @@ -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(()) + } +} diff --git a/src/exec/rmbg2_0.rs b/src/exec/rmbg2_0.rs new file mode 100644 index 0000000..a93128a --- /dev/null +++ b/src/exec/rmbg2_0.rs @@ -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(()) + } +} diff --git a/src/exec/voxcpm.rs b/src/exec/voxcpm.rs new file mode 100644 index 0000000..7a2e2d3 --- /dev/null +++ b/src/exec/voxcpm.rs @@ -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(()) + } +} diff --git a/src/exec/voxcpm1_5.rs b/src/exec/voxcpm1_5.rs new file mode 100644 index 0000000..b6c7d6b --- /dev/null +++ b/src/exec/voxcpm1_5.rs @@ -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(()) + } +} diff --git a/src/lib.rs b/src/lib.rs index 49d7fff..52d096c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,4 +1,5 @@ pub mod chat_template; +pub mod exec; pub mod models; pub mod position_embed; pub mod tokenizer; diff --git a/src/main.rs b/src/main.rs index 0a1cb22..1cbc4be 100644 --- a/src/main.rs +++ b/src/main.rs @@ -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, + /// Service listen port #[arg(short, long)] - model: WhichModel, + port: Option, + /// Model type (required for backward compatibility) + #[arg(short, long)] + model: Option, + + /// Local model weight path #[arg(long)] weight_path: Option, + /// Model download save directory #[arg(long)] save_dir: Option, + /// Download retry count + #[arg(long)] + download_retries: Option, + + #[command(subcommand)] + command: Option, +} + +#[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, + + /// Model download save directory + #[arg(long)] + save_dir: Option, + + /// Download retry count #[arg(long)] download_retries: Option, } -// 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, + + /// Download retry count + #[arg(long)] + download_retries: Option, +} + +/// 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, + + /// Output file path (optional) + #[arg(short, long)] + output: Option, + + /// 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!"); -// } diff --git a/src/models/mod.rs b/src/models/mod.rs index ceedfe4..0299553 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -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, } diff --git a/src/models/qwen3_asr/config.rs b/src/models/qwen3_asr/config.rs new file mode 100644 index 0000000..b023260 --- /dev/null +++ b/src/models/qwen3_asr/config.rs @@ -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, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3ASRGenerationConfig { + pub do_sample: bool, + pub eos_token_id: Vec, + pub pad_token_id: usize, + pub temperature: f32, +} \ No newline at end of file diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs new file mode 100644 index 0000000..f1e12ab --- /dev/null +++ b/src/models/qwen3_asr/generate.rs @@ -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) -> Result { + 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 { +// 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> +// + 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))) +// } +// } diff --git a/src/models/qwen3_asr/mod.rs b/src/models/qwen3_asr/mod.rs new file mode 100644 index 0000000..8b1baf7 --- /dev/null +++ b/src/models/qwen3_asr/mod.rs @@ -0,0 +1,4 @@ +pub mod config; +pub mod generate; +pub mod model; +pub mod processor; diff --git a/src/models/qwen3_asr/model.rs b/src/models/qwen3_asr/model.rs new file mode 100644 index 0000000..d3f971f --- /dev/null +++ b/src/models/qwen3_asr/model.rs @@ -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 { + 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 { + 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 { + 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 { + 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 { + 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, + 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 { + 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 { + 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 { + 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, + tp_encoders: Vec, + 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 { + 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 { + 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, + normalize_before: bool, +} + +impl AdaptorEncoderLayer { + pub fn new( + vb: VarBuilder, + llm_dim: usize, + n_head: usize, + normalize_before: bool, + concat_after: bool, + ) -> Result { + 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 { + 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, +} + +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 { + 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 { + 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 { + 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 { + 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::()?; + 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(); + } +} diff --git a/src/models/qwen3_asr/processor.rs b/src/models/qwen3_asr/processor.rs new file mode 100644 index 0000000..508ac6b --- /dev/null +++ b/src/models/qwen3_asr/processor.rs @@ -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, + max_asr_input_seconds: f32, +} + +impl Qwen3AsrProcessor { + pub fn new(device: &Device) -> Result { + let support_language: Vec = 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> { + 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 {}''", 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, +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index d8155c5..a1bf32a 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -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 { + 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) +} diff --git a/tests/test_qwen3_asr.rs b/tests/test_qwen3_asr.rs new file mode 100644 index 0000000..acb7937 --- /dev/null +++ b/tests/test_qwen3_asr.rs @@ -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(()) +// } diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index 9c81761..7f0f733 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -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