初始化 Qwen-Asr 本地仓库

main
Bifang 2026-09-28 17:00:53 +08:00
commit dde3e12476
192 changed files with 56639 additions and 0 deletions

70
.dockerignore 100644
View File

@ -0,0 +1,70 @@
# Git文件
.git
.gitignore
.gitmodules
# Python缓存
__pycache__
*.pyc
*.pyo
*.pyd
.Python
*.so
# 虚拟环境
venv
env
ENV
.venv
# IDE文件
.vscode
.idea
*.swp
*.swo
# 系统文件
.DS_Store
Thumbs.db
# 日志文件
*.log
logs/
# 临时文件
temp/
tmp/
*.tmp
build-file/
# 模型文件(这些将通过volume挂载)
tts/third_party/CosyVoice/pretrained_models/
models/
*.pt
*.pth
*.bin
*.safetensors
*.tar.gz
# 测试文件
tests/
*.test
coverage.*
# 文档
README.md
*.md
docs/
# 配置文件(可能包含敏感信息)
.env
.env.*
# 本地音视频样本
*.m4a
*.mp4
# 其他
.pytest_cache
.coverage
node_modules

150
.env.example 100644
View File

@ -0,0 +1,150 @@
# Qwen3-ASR 环境变量覆盖示例
# 复制为 .env 后,只取消你需要修改的项的注释即可。
# -----------------------------------------------------------------------------
# 仅在需要 API Key 鉴权时开启。
# -----------------------------------------------------------------------------
# API_KEY=your_api_key_here
# -----------------------------------------------------------------------------
# Docker / 部署入口配置。
# -----------------------------------------------------------------------------
# NGINX_PORT:宿主机暴露的 Nginx 端口。
# NGINX_PORT=17003
# ASR_IMAGE:离线包或自定义部署时使用的镜像标签。
# ASR_IMAGE=unis/qwen3-asr:gpu-latest
# ASR_VISIBLE_DEVICES:统一可见卡编号配置,多个用逗号分隔;程序会按当前 accelerator 自动映射到底层变量。
# ASR_VISIBLE_DEVICES=0
# METAX_DRI_DEVICE:沐曦 / MuXi 需要映射的 DRI 设备路径。
# METAX_DRI_DEVICE=/dev/dri
# METAX_MXSMI_PATH:沐曦 / MuXi 的 mx-smi 工具路径。
# METAX_MXSMI_PATH=/opt/mxdriver/bin/mx-smi
# METAX_PRIVILEGED:沐曦 / MuXi 容器是否以 privileged 方式运行。
# METAX_PRIVILEGED=false
# ILUVATAR_USR_SRC:天数 / Iluvatar 容器内 usr/src 挂载路径。
# ILUVATAR_USR_SRC=/usr/src
# ILUVATAR_LIB_MODULES:天数 / Iluvatar 容器内模块目录挂载路径。
# ILUVATAR_LIB_MODULES=/lib/modules
# ILUVATAR_DEV:天数 / Iluvatar 设备目录挂载路径。
# ILUVATAR_DEV=/dev
# ILUVATAR_HOME:天数 / Iluvatar 主目录挂载路径。
# ILUVATAR_HOME=/home
# ILUVATAR_DATA:天数 / Iluvatar 数据目录挂载路径。
# ILUVATAR_DATA=/data
# MTHREADS_DEV:摩尔线程 / Moore Threads 设备目录挂载路径。
# MTHREADS_DEV=/dev
# MTHREADS_USR_SRC:摩尔线程 / Moore Threads 容器内 usr/src 挂载路径。
# MTHREADS_USR_SRC=/usr/src
# MTHREADS_LIB_MODULES:摩尔线程 / Moore Threads 容器内模块目录挂载路径。
# MTHREADS_LIB_MODULES=/lib/modules
# MTHREADS_HOME:摩尔线程 / Moore Threads 主目录挂载路径。
# MTHREADS_HOME=/home
# MTHREADS_DATA:摩尔线程 / Moore Threads 数据目录挂载路径。
# MTHREADS_DATA=/data
# MODEL_STORAGE_DIR:模型宿主机挂载目录。
# MODEL_STORAGE_DIR=/opt/dep/asr/models
# DATA_STORAGE_DIR:业务数据宿主机挂载目录。
# DATA_STORAGE_DIR=/opt/dep/asr/data
# NGINX_RATE_LIMIT_RPS:Nginx 全局限流,单位为每秒请求数。
# NGINX_RATE_LIMIT_RPS=0
# NGINX_RATE_LIMIT_BURST:Nginx 全局突发限流额度,0 表示自动按 RPS 处理。
# NGINX_RATE_LIMIT_BURST=0
# ASR_DEPLOY_TOPOLOGY:部署拓扑,isolated=每卡一个实例,sharded=单实例多卡分片,auto=优先 sharded 失败回退 isolated。
# ASR_DEPLOY_TOPOLOGY=isolated
# -----------------------------------------------------------------------------
# 运行时模型选择。
# 留空 QWEN3_ASR_MODEL 时会自动选择。
# 支持值:qwen3-asr-0.6b、qwen3-asr-1.7b
# -----------------------------------------------------------------------------
# ACCELERATOR:加速后端类型,常见值有 auto、cpu、nvidia、metax、iluvatar、mthreads。
# ACCELERATOR=auto
# DEVICE:具体运行设备,常见值有 auto、cpu、cuda:0。
# DEVICE=auto
# QWEN3_ASR_MODEL:手动指定离线识别模型,留空则自动挑选。
# QWEN3_ASR_MODEL=
# -----------------------------------------------------------------------------
# 模型下载 / 缓存行为。
# 在 Docker Compose 中,默认宿主机挂载目录是 /opt/dep/asr/models。
# 实际模型目录通常位于:
# /opt/dep/asr/models/Qwen
# /opt/dep/asr/models/iic
# /opt/dep/asr/models/damo
# -----------------------------------------------------------------------------
# MODELS_DIR:项目主模型目录。
# MODELS_DIR=/opt/dep/asr/models
# MODELSCOPE_CACHE:ModelScope 缓存根目录,下面会生成 models/{publisher}/{model}。
# MODELSCOPE_CACHE=/opt/dep/asr
# MODELSCOPE_PATH:ModelScope 实际模型目录,建议保持在 models 目录下。
# MODELSCOPE_PATH=/opt/dep/asr/models
# -----------------------------------------------------------------------------
# 说话人注册 / pgvector 数据库配置。
# 表结构与 Model-Test-New 保持一致:speakers(id, name, user_id, embedding)。
# -----------------------------------------------------------------------------
# SPEAKER_DB_ENABLED:是否启用说话人库与 pgvector 检索。
# SPEAKER_DB_ENABLED=true
# DB_HOST:PostgreSQL 主机地址。
# DB_HOST=127.0.0.1
# DB_PORT:PostgreSQL 端口。
# DB_PORT=5432
# DB_USER:PostgreSQL 用户名。
# DB_USER=postgres
# DB_PASSWORD:PostgreSQL 密码。
# DB_PASSWORD=postgres
# DB_NAME:PostgreSQL 数据库名。
# DB_NAME=asr_db
# SV_MODEL:说话人识别模型。
# SV_MODEL=iic/speech_campplus_sv_zh-cn_16k-common
# SV_THRESHOLD:说话人相似度阈值,越大越严格。
# SV_THRESHOLD=0.6
# TEMP_DIR:临时文件目录。
# TEMP_DIR=/opt/dep/asr/data/temp
# LOG_FILE:日志文件路径。
# LOG_FILE=/opt/dep/asr/data/logs/qwen3-asr.log
# TASK_STATE_DIR:任务状态持久化目录。
# TASK_STATE_DIR=/opt/dep/asr/data/tasks
# TASK_RETENTION_HOURS:任务结果保留时间,单位小时。
# TASK_RETENTION_HOURS=24
# -----------------------------------------------------------------------------
# CPU Rust 后端覆盖配置。
# 一般不需要改;自动检测会检查 vendor/qwenasr/target/{release,debug}。
# -----------------------------------------------------------------------------
# QWENASR_LIBRARY_PATH:手动指定 libqwen_asr.so 的绝对路径。
# QWENASR_LIBRARY_PATH=/absolute/path/to/libqwen_asr.so
# -----------------------------------------------------------------------------
# 调优参数。除非你在做特定瓶颈测试,否则建议保持默认。
# -----------------------------------------------------------------------------
# ASR_BATCH_SIZE:批处理大小,表示一次并行推理的片段数。
# ASR_BATCH_SIZE=4
# ASR_ENABLE_WORD_TIMESTAMPS:是否启用字词级时间戳。
# ASR_ENABLE_WORD_TIMESTAMPS=false
# MAX_SEGMENT_SEC:离线 ASR 单段最大时长,单位秒。
# MAX_SEGMENT_SEC=60
# QWEN_RUST_CPU_WORKERS:Qwen Rust CPU 后端 worker 数。
# QWEN_RUST_CPU_WORKERS=4
# QWEN_RUST_ASR_CONCURRENCY:Qwen Rust ASR 并发数,0 通常表示自动。
# QWEN_RUST_ASR_CONCURRENCY=0
# QWEN_RUST_ALIGN_CONCURRENCY:Qwen Rust 对齐并发数,0 通常表示自动。
# QWEN_RUST_ALIGN_CONCURRENCY=0
# QWEN_GPU_MEMORY_UTILIZATION:GPU 显存使用比例。
# QWEN_GPU_MEMORY_UTILIZATION=0.9
# QWEN_VLLM_ENFORCE_EAGER:是否强制 vLLM eager 执行;国产卡/稳定优先建议 true,NVIDIA 性能测试可设 false。
# QWEN_VLLM_ENFORCE_EAGER=true
# QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION:forced aligner 预留的显存比例。
# QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION=0.15
# ASR_ENABLE_NEARFIELD_FILTER:是否启用近场/远场过滤。
# ASR_ENABLE_NEARFIELD_FILTER=true
# ASR_NEARFIELD_RMS_THRESHOLD:近场判断的 RMS 能量阈值。
# ASR_NEARFIELD_RMS_THRESHOLD=0.01
# -----------------------------------------------------------------------------
# 仅用于本地开发调试。
# -----------------------------------------------------------------------------
# LOG_LEVEL:日志级别。
# LOG_LEVEL=INFO
# FUNASR_STARTUP_UI:FunASR 启动界面模式。
# FUNASR_STARTUP_UI=auto

48
.gitignore vendored 100644
View File

@ -0,0 +1,48 @@
# Environment files (contain sensitive information)
.env
*.env.local
*.env.*.local
# Build and deployment files
build-file/
# Downloaded models directory
/models/
# Large binary packages
*.tar.gz
*.tar
*.zip
# Python cache
__pycache__/
*.py[cod]
*$py.class
*.so
# Model files (downloaded)
*.pt
*.pth
*.onnx
*.safetensors
*.bin
# Log files
*.log
logs/
# OS files
.DS_Store
Thumbs.db
# IDE
.idea/
.vscode/
*.swp
*.swo
# Temporary files
*.tmp
*.temp
.cache/
.code-review-graph/

View File

@ -0,0 +1,36 @@
# Code Review Graph MCP
Generated from the `agent-tool-basic` template.
## Contributions
- Agent tool `echo_text`
## Develop
1. Open the Plugins page and use **Load development plugin**, pointing at this
directory. PI-Desktop reloads the plugin whenever you save a file here.
2. Verify the contributions from the command palette.
3. Validate and package:
```bash
pnpm pi-plugin check .
pnpm pi-plugin pack .
# writes dist/crg-mcp-0.1.0.piplug
```
Install the resulting `.piplug` from the Plugins page to test it the way a
user would.
### Panel top drag band
PI-Desktop reserves exactly a transparent 46px frameless drag band above panel
content and renders a minimal fixed three-button window-control capsule in its
top-right corner. Normal-flow content is offset automatically. The panel title,
toolbar, and every other visible surface belong to the plugin. Development
panels show a reminder that the top 46px is not clickable outside the capsule.
For `position: fixed` or `position: sticky` content, use
`top: var(--pi-plugin-titlebar-height, 46px)` and account for the same value
in viewport-height calculations. Add `-webkit-app-region: drag` to a
plugin-owned toolbar when it should move the window, and
`-webkit-app-region: no-drag` to controls inside it.

View File

@ -0,0 +1,57 @@
@echo off
chcp 65001 >nul
setlocal
rem ---------------------------------------------------------------------------
rem code-review-graph MCP launcher for PI-Desktop.
rem
rem The MCP host spawns this with a minimal environment (PATH, SystemRoot,
rem windir, TEMP, TMP, LANG plus the manifest's env block) and cwd = plugin
rem directory, with stdin/stdout used for JSON-RPC. Never write to stdout.
rem ---------------------------------------------------------------------------
if defined CRG_HOME goto have_home
set "CRG_HOME=D:\github-project\code-review-graph"
:have_home
rem code_review_graph/constants.py calls Path.home() at import time; without a
rem user profile the server dies with "Could not determine home directory.".
if defined USERPROFILE goto have_profile
set "USERPROFILE=C:\Users\%USERNAME%"
:have_profile
if defined HOMEDRIVE goto have_hd
set "HOMEDRIVE=C:"
:have_hd
if defined HOMEPATH goto have_hp
set "HOMEPATH=\Users\%USERNAME%"
:have_hp
if defined APPDATA set "APPDATA=%USERPROFILE%\AppData\Roaming"
if defined LOCALAPPDATA set "LOCALAPPDATA=%USERPROFILE%\AppData\Local"
set "PYTHONUTF8=1"
set "PYTHONIOENCODING=utf-8"
rem The editable install's .pth points at D:\github_project\... (underscore)
rem while the checkout lives at D:\github-project\... (hyphen), so the package
rem is only importable with the checkout explicitly on sys.path.
set "PYTHONPATH=%CRG_HOME%"
rem The venv's code-review-graph.exe is a broken uv trampoline
rem ("failed to canonicalize script path"), so prefer a working interpreter.
set "CRG_PY=%CRG_HOME%\.venv\Scripts\python.exe"
if exist "%CRG_PY%" goto py_ready
set "CRG_PY=python"
:py_ready
rem Upstream ships 30 tools; keep the surface small unless overridden.
if defined CRG_TOOLS goto tools_ready
set "CRG_TOOLS=build_or_update_graph_tool,run_postprocess_tool,get_minimal_context_tool,get_review_context_tool,get_impact_radius_tool,query_graph_tool,semantic_search_nodes_tool,detect_changes_tool,list_graph_stats_tool,get_affected_flows_tool"
:tools_ready
if defined CRG_REPO goto repo_ready
set "CRG_REPO=%CRG_HOME%"
:repo_ready
"%CRG_PY%" -m code_review_graph serve --repo "%CRG_REPO%"
endlocal

View File

@ -0,0 +1,16 @@
/**
* crg-mcp — thin loader for the code-review-graph MCP bridge.
*
* All wiring lives in manifest.json under `contributes.mcpServers`: the host
* spawns `crg.cmd`, speaks MCP over its stdio, and publishes each upstream tool
* as `plugin_crg_mcp_crg_<tool>`. The host owns the client, the retry on the
* next call, and the tool lifecycle, so there is nothing to register here.
*/
async function onLoad() {
// Intentionally empty: the MCP client is owned by the host.
}
async function onUnload() {}
module.exports = { onLoad, onUnload };

View File

@ -0,0 +1,33 @@
{
"schemaVersion": 1,
"id": "crg-mcp",
"name": "Code Review Graph MCP",
"version": "1.0.0",
"description": "Bridges the locally deployed code-review-graph MCP server into the agent as native tools.",
"main": "main.js",
"contributes": {
"mcpServers": [
{
"id": "crg",
"label": "Code Review Graph",
"transport": "stdio",
"command": "crg.cmd",
"env": {
"PYTHONPATH": "D:\\github-project\\code-review-graph",
"USERPROFILE": "C:\\Users\\admin",
"HOMEDRIVE": "C:",
"HOMEPATH": "\\Users\\admin"
}
}
]
},
"permissions": [
"mcp.server.local"
],
"engines": {
"piDesktop": ">=0.1.0"
},
"activationEvents": [
"onStartup"
]
}

View File

@ -0,0 +1,96 @@
<#
Self-check for the crg-mcp plugin.
Drives crg.cmd exactly the way the PI-Desktop MCP host does — `cmd /c crg.cmd`
with piped stdio and cwd = this folder — then reports the MCP handshake, the
tool list, and one real tool call. Run it from PowerShell:
powershell -NoProfile -ExecutionPolicy Bypass -File .\selfcheck.ps1
#>
$ErrorActionPreference = 'Stop'
$dir = Split-Path -Parent $MyInvocation.MyCommand.Path
# Mirror the host's minimal environment plus the manifest's env block.
foreach ($k in 'PATH', 'SystemRoot', 'TEMP', 'TMP') {
if (-not (Test-Path "Env:$k")) { Write-Warning "missing $k in ambient env" }
}
$psi = New-Object System.Diagnostics.ProcessStartInfo
$psi.FileName = 'cmd.exe'
$psi.Arguments = '/c crg.cmd'
$psi.WorkingDirectory = $dir
$psi.RedirectStandardInput = $true
$psi.RedirectStandardOutput = $true
$psi.RedirectStandardError = $true
$psi.UseShellExecute = $false
$psi.StandardOutputEncoding = [System.Text.Encoding]::UTF8
$proc = [System.Diagnostics.Process]::Start($psi)
function Send($obj) {
$proc.StandardInput.WriteLine(($obj | ConvertTo-Json -Compress -Depth 8))
$proc.StandardInput.Flush()
}
function ReadLine([int]$waitSeconds = 25) {
$task = $proc.StandardOutput.ReadLineAsync()
if ($task.Wait([TimeSpan]::FromSeconds($waitSeconds))) { return $task.Result }
return $null
}
Send @{
jsonrpc = '2.0'; id = 1; method = 'initialize'
params = @{
protocolVersion = '2025-06-18'
capabilities = @{}
clientInfo = @{ name = 'crg-selfcheck'; version = '1' }
}
}
$initLine = ReadLine 40
if (-not $initLine) {
Write-Host 'HANDSHAKE FAILED: no stdout from crg.cmd' -ForegroundColor Red
Write-Host '--- stderr ---'
Write-Host $proc.StandardError.ReadToEnd()
try { $proc.Kill() } catch { }
exit 1
}
$init = $initLine | ConvertFrom-Json
Write-Host ("HANDSHAKE OK server={0} {1}" -f $init.result.serverInfo.name, $init.result.serverInfo.version) -ForegroundColor Green
Send @{ jsonrpc = '2.0'; method = 'notifications/initialized'; params = @{} }
Send @{ jsonrpc = '2.0'; id = 2; method = 'tools/list'; params = @{} }
$toolsLine = ReadLine
if (-not $toolsLine) {
Write-Host 'tools/list returned nothing' -ForegroundColor Red
try { $proc.Kill() } catch { }
exit 1
}
$tools = ($toolsLine | ConvertFrom-Json).result.tools
Write-Host ("TOOLS: {0}" -f $tools.Count) -ForegroundColor Green
foreach ($t in $tools) { Write-Host (" plugin_crg_mcp_crg_{0}" -f $t.name) }
Send @{
jsonrpc = '2.0'; id = 3; method = 'tools/call'
params = @{ name = 'list_graph_stats_tool'; arguments = @{} }
}
$callLine = ReadLine 40
if ($callLine) {
$call = $callLine | ConvertFrom-Json
if ($call.result) {
$text = $call.result.content[0].text
Write-Host 'TOOL CALL OK' -ForegroundColor Green
Write-Host (' ' + ($text -split "`n")[0..3] -join ' | ')
}
else {
Write-Host ("TOOL CALL ERROR: {0}" -f ($call | ConvertTo-Json -Compress -Depth 6)) -ForegroundColor Red
}
}
else {
Write-Host 'TOOL CALL: no response' -ForegroundColor Red
}
try { $proc.Kill() } catch { }

73
Dockerfile.cpu 100644
View File

@ -0,0 +1,73 @@
FROM python:3.10-slim AS runtime
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
HF_HUB_DISABLE_SYMLINKS_WARNING=1 \
HF_HUB_DISABLE_PROGRESS_BARS=1 \
OPENBLAS_NUM_THREADS=1 \
OMP_NUM_THREADS=1 \
GOTO_NUM_THREADS=1 \
QWENASR_LIBRARY_PATH=/opt/qwenasr/lib/libqwen_asr.so
# Install system packages required for audio processing
RUN apt-get update && apt-get install -y --no-install-recommends \
ffmpeg \
sox \
libsox-dev \
libsndfile1 \
libopenblas-dev \
nginx \
build-essential \
curl
WORKDIR /app
ARG TARGETARCH
ARG QWENASR_RUST_TARGET_CPU=x86-64-v2
COPY vendor/qwenasr /tmp/qwenasr
# Install Rust compiler (required for sudachipy on ARM64)
RUN curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y
ENV PATH="/root/.cargo/bin:${PATH}"
# Build QwenASR Rust shared library for CPU Qwen3-ASR inference. The default
# amd64 target stays portable across common AVX2-era servers; set
# QWENASR_RUST_TARGET_CPU=native only for self-built, host-specific images.
RUN if [ "$TARGETARCH" = "amd64" ]; then export RUSTFLAGS="-C target-cpu=${QWENASR_RUST_TARGET_CPU}"; fi \
&& cargo build --release -p qwen-asr --features ffi --manifest-path /tmp/qwenasr/Cargo.toml \
&& mkdir -p /opt/qwenasr/lib \
&& cp /tmp/qwenasr/target/release/libqwen_asr.so /opt/qwenasr/lib/libqwen_asr.so
# Install Python dependencies (CPU mode)
COPY environments/cpu/pyproject.toml /app/environments/cpu/pyproject.toml
RUN python - <<'PY' > /tmp/qwen3-asr-cpu-reqs.txt
import tomllib
from pathlib import Path
data = tomllib.loads(Path("/app/environments/cpu/pyproject.toml").read_text())
for dep in data["project"]["dependencies"]:
print(dep)
print("asyncpg==0.31.0")
PY
RUN pip install --no-cache-dir -r /tmp/qwen3-asr-cpu-reqs.txt && \
rm -f /tmp/qwen3-asr-cpu-reqs.txt
# Clean build tools and cache to reduce image size
RUN apt remove -y build-essential curl && apt autoremove -y \
&& rm -rf /tmp/qwenasr \
&& rm -rf /root/.cargo /root/.rustup \
&& apt-get clean && rm -rf /var/lib/apt/lists/*
# Copy application code
COPY . .
# Create runtime directories
RUN mkdir -p /app/data/temp /app/data/logs /app/data/tasks \
&& chmod +x start.py /app/scripts/docker/entrypoint.sh
EXPOSE 8000
ENTRYPOINT ["/app/scripts/docker/entrypoint.sh"]
CMD ["python", "start.py"]

101
Dockerfile.gpu 100644
View File

@ -0,0 +1,101 @@
ARG PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda13.0-cudnn9-runtime
FROM ${PYTORCH_BASE_IMAGE}
ARG CUDA_NVCC_PACKAGE=cuda-nvcc-13-0
ARG PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu130
ARG TORCH_VERSION=2.11.0
ARG TORCHAUDIO_VERSION=2.11.0
ARG TORCHVISION_VERSION=0.26.0
ARG VLLM_VERSION=0.20.0
ARG TORCH_CUDA_ARCH_LIST=12.0+PTX
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
HF_HUB_DISABLE_SYMLINKS_WARNING=1 \
HF_HUB_DISABLE_PROGRESS_BARS=1 \
PIP_BREAK_SYSTEM_PACKAGES=1 \
TORCH_CUDA_ARCH_LIST=${TORCH_CUDA_ARCH_LIST}
# Add NVIDIA apt repository for CUDA packages
RUN apt-get update && apt-get install -y --no-install-recommends \
ca-certificates \
gnupg \
wget \
&& wget -qO - https://developer.download.nvidia.com/compute/cuda/repos/ubuntu2404/x86_64/3bf863cc.pub | apt-key add - \
&& echo "deb https://developer.download.nvidia.com/compute/cuda/repos/ubuntu2404/x86_64 /" > /etc/apt/sources.list.d/cuda.list \
&& apt-get update
# Install system packages required for audio processing
# Note: nvcc and build-essential are required for FlashInfer JIT compilation.
RUN apt-get install -y --no-install-recommends \
ffmpeg \
sox \
libsox-dev \
libsndfile1 \
nginx \
build-essential \
${CUDA_NVCC_PACKAGE}
WORKDIR /app
# Install Python dependencies directly into the image runtime Python.
# The image defaults to CUDA 13.0/cu130 for Blackwell-capable GPUs. Developers
# can rebuild with a different PyTorch CUDA backend by overriding:
# PYTORCH_BASE_IMAGE, PYTORCH_CUDA_INDEX, CUDA_NVCC_PACKAGE, TORCH_CUDA_ARCH_LIST.
COPY pyproject.toml /app/
RUN python - <<'PY' > /tmp/qwen3-asr-gpu-reqs.txt
import tomllib
from pathlib import Path
data = tomllib.loads(Path("/app/pyproject.toml").read_text())
for dep in data["project"]["dependencies"]:
if dep.startswith("torch=="):
continue
if dep.startswith("torchaudio=="):
continue
if dep.startswith("torchvision=="):
continue
if dep.startswith("vllm=="):
continue
print(dep)
PY
RUN pip install --no-cache-dir -r /tmp/qwen3-asr-gpu-reqs.txt && \
pip install --no-cache-dir \
--index-url "${PYTORCH_CUDA_INDEX}" \
--extra-index-url https://pypi.org/simple \
"torch==${TORCH_VERSION}" \
"torchaudio==${TORCHAUDIO_VERSION}" \
"torchvision==${TORCHVISION_VERSION}" && \
pip install --no-cache-dir "vllm==${VLLM_VERSION}" && \
rm -f /tmp/qwen3-asr-gpu-reqs.txt
# Fail the image build early if the runtime dependency chain is inconsistent.
RUN python - <<'PY'
import torch
import torchaudio
import transformers
import vllm
from transformers import PreTrainedModel
print("torch", torch.__version__)
print("torchaudio", torchaudio.__version__)
print("transformers", transformers.__version__)
print("vllm", vllm.__version__)
print("PreTrainedModel", PreTrainedModel)
PY
# Clean apt cache but keep build-essential and nvcc for FlashInfer JIT compilation.
RUN apt-get clean && rm -rf /var/lib/apt/lists/*
# Copy application code
COPY . .
# Create runtime directories
RUN mkdir -p /app/data/temp /app/data/logs /app/data/tasks \
&& chmod +x start.py /app/scripts/docker/entrypoint.sh
EXPOSE 8000
ENTRYPOINT ["/app/scripts/docker/entrypoint.sh"]
CMD ["python", "start.py"]

View File

@ -0,0 +1,51 @@
ARG ILUVATAR_BASE_IMAGE=registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
FROM ${ILUVATAR_BASE_IMAGE}
ARG PYTHON_BIN=python3
ARG INSTALL_SYSTEM_PACKAGES=false
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
HF_HUB_DISABLE_SYMLINKS_WARNING=1 \
HF_HUB_DISABLE_PROGRESS_BARS=1 \
PIP_BREAK_SYSTEM_PACKAGES=1 \
ACCELERATOR=iluvatar \
DEVICE=auto \
APP_PYTHON_BIN=${PYTHON_BIN}
# Iluvatar production images should be built on top of the official Iluvatar
# vLLM image. The base image owns the IX runtime, PyTorch, vLLM, and kernels;
# this layer only adds generic project/runtime dependencies and app code.
RUN if [ "${INSTALL_SYSTEM_PACKAGES}" = "true" ]; then \
apt-get update && apt-get install -y --no-install-recommends \
python3 python3-pip python3-venv ffmpeg sox libsox-dev libsndfile1 nginx \
&& apt-get clean && rm -rf /var/lib/apt/lists/*; \
fi
WORKDIR /app
COPY environments/iluvatar/requirements.txt /tmp/qwen3-asr-iluvatar-requirements.txt
RUN ${APP_PYTHON_BIN} -m pip install --no-cache-dir -r /tmp/qwen3-asr-iluvatar-requirements.txt && \
rm -f /tmp/qwen3-asr-iluvatar-requirements.txt
RUN ${APP_PYTHON_BIN} - <<'PY'
import importlib
import torch
print("torch", torch.__version__)
print("torch.cuda.is_available", torch.cuda.is_available())
for name in ("vllm",):
module = importlib.import_module(name)
print(name, getattr(module, "__version__", "unknown"))
PY
COPY . .
RUN mkdir -p /app/data/temp /app/data/logs /app/data/tasks \
&& chmod +x start.py /app/scripts/docker/entrypoint.sh
EXPOSE 8000
ENTRYPOINT ["/app/scripts/docker/entrypoint.sh"]
CMD ["python3", "start.py"]

55
Dockerfile.metax 100644
View File

@ -0,0 +1,55 @@
ARG METAX_BASE_IMAGE
FROM ${METAX_BASE_IMAGE}
ARG PYTHON_BIN=/opt/conda/bin/python
ARG INSTALL_SYSTEM_PACKAGES=true
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
HF_HUB_DISABLE_SYMLINKS_WARNING=1 \
HF_HUB_DISABLE_PROGRESS_BARS=1 \
PIP_BREAK_SYSTEM_PACKAGES=1 \
ACCELERATOR=metax \
DEVICE=auto \
APP_PYTHON_BIN=${PYTHON_BIN}
# MetaX production images should be built on top of an official MetaX vLLM
# image that already contains MACA, PyTorch, vLLM, and their compiled kernels.
# We only add generic project/runtime dependencies and the application code.
RUN if [ "${INSTALL_SYSTEM_PACKAGES}" = "true" ]; then \
apt-get update && apt-get install -y --no-install-recommends \
python3 python3-pip python3-venv ffmpeg sox libsox-dev libsndfile1 nginx \
&& apt-get clean && rm -rf /var/lib/apt/lists/*; \
fi
WORKDIR /app
COPY environments/metax/requirements.txt /tmp/qwen3-asr-metax-requirements.txt
RUN ${APP_PYTHON_BIN} -m pip install \
--no-cache-dir \
--no-warn-conflicts \
--root-user-action=ignore \
-r /tmp/qwen3-asr-metax-requirements.txt && \
rm -f /tmp/qwen3-asr-metax-requirements.txt
RUN ${APP_PYTHON_BIN} - <<'PY'
import importlib
import torch
print("torch", torch.__version__)
print("torch.cuda.is_available", torch.cuda.is_available())
for name in ("vllm",):
module = importlib.import_module(name)
print(name, getattr(module, "__version__", "unknown"))
PY
COPY . .
RUN mkdir -p /app/data/temp /app/data/logs /app/data/tasks \
&& chmod +x start.py /app/scripts/docker/entrypoint.sh
EXPOSE 8000
ENTRYPOINT ["/app/scripts/docker/entrypoint.sh"]
CMD ["/opt/conda/bin/python", "start.py"]

View File

@ -0,0 +1,40 @@
ARG MTHREADS_BASE_IMAGE=registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
FROM ${MTHREADS_BASE_IMAGE}
ARG PYTHON_BIN=python3
ARG INSTALL_SYSTEM_PACKAGES=true
ENV DEBIAN_FRONTEND=noninteractive \
PYTHONUNBUFFERED=1 \
PYTHONDONTWRITEBYTECODE=1 \
HF_HUB_DISABLE_SYMLINKS_WARNING=1 \
HF_HUB_DISABLE_PROGRESS_BARS=1 \
PIP_BREAK_SYSTEM_PACKAGES=1 \
ACCELERATOR=mthreads \
DEVICE=auto \
APP_PYTHON_BIN=${PYTHON_BIN}
# Moore Threads production images should be built on top of the official MUSA
# vLLM image. Keep the vendor Python/runtime stack intact and only add the
# extra system tools needed by the application.
RUN if [ "${INSTALL_SYSTEM_PACKAGES}" = "true" ]; then \
apt-get update && apt-get install -y --no-install-recommends \
ffmpeg sox libsox-dev libsndfile1 nginx \
&& apt-get clean && rm -rf /var/lib/apt/lists/*; \
fi
WORKDIR /app
COPY environments/mthreads/requirements.txt /tmp/qwen3-asr-mthreads-requirements.txt
RUN ${APP_PYTHON_BIN} -m pip install --no-cache-dir -r /tmp/qwen3-asr-mthreads-requirements.txt && \
rm -f /tmp/qwen3-asr-mthreads-requirements.txt
COPY . .
RUN mkdir -p /app/data/temp /app/data/logs /app/data/tasks \
&& chmod +x start.py /app/scripts/docker/entrypoint.sh
EXPOSE 8000
ENTRYPOINT ["/app/scripts/docker/entrypoint.sh"]
CMD ["python3", "start.py"]

610
README.md 100644
View File

@ -0,0 +1,610 @@
<div align="center">
<h1>Qwen3-ASR</h1>
<h3>Ready-to-use Local Speech Recognition API Service</h3>
Speech recognition API service centered on [Qwen3-ASR](https://github.com/QwenLM/Qwen3-ASR), with NVIDIA CUDA vLLM, MetaX/MuXi MACA vLLM, and CPU Rust backends, OpenAI API compatibility, Alibaba Cloud Speech API compatibility, and a Paraformer realtime websocket capability.
[简体中文](./docs/README_zh.md)
---
![Static Badge](https://img.shields.io/badge/Python-3.10+-blue?logo=python)
![Static Badge](https://img.shields.io/badge/Torch-2.11.0-%23EE4C2C?logo=pytorch&logoColor=white)
![Static Badge](https://img.shields.io/badge/CUDA-13.0_default-%2376B900?logo=nvidia&logoColor=white)
</div>
## Live Demo Site
- **Web Demo**: https://asr.vect.one
## Demo
[![Demo](./demo/demo.png)](https://media.cdn.vect.one/qwenasr_client_demo.mp4)
## Release 1.0.1
> `v1.0.1` is the current patch release. `v1.0.0` introduced a large breaking refactor relative to the earlier `main` branch.
> If you are upgrading from `main`, read the release notes before reusing old deployment assumptions.
>
> Key breaking changes:
> - Python dependency management is now `uv`-based (`pyproject.toml` + `uv.lock`); `requirements*.txt` are gone
> - Runtime stack changed to `NVIDIA/MetaX GPU -> vLLM`, `CPU/macOS -> vendored QwenASR Rust`
> - `MLX` / Apple Silicon GPU path has been removed; `mps` is normalized to `cpu`
> - macOS / Apple Silicon now defaults to `qwen3-asr-0.6b`; set `QWEN3_ASR_MODEL` to override it
> - `ENABLED_MODELS` has been removed
## Features
- **Hybrid Runtime Stack** - Uses auto-selected Qwen3-ASR for offline inference and Paraformer realtime for websocket streaming
- **Speaker Diarization** - Automatic multi-speaker identification using CAM++ model
- **OpenAI API Compatible** - Supports `/v1/audio/transcriptions` endpoint, works with OpenAI SDK
- **Alibaba Cloud API Compatible** - Supports Alibaba Cloud Speech RESTful API and WebSocket streaming protocol
- **WebSocket Streaming** - Real-time streaming speech recognition with low latency
- **Smart Far-Field Filtering** - Automatically filters far-field sounds and ambient noise in streaming ASR
- **Intelligent Audio Segmentation** - VAD-based greedy merge algorithm for automatic long audio splitting
- **GPU Batch Processing** - Batch inference support, 2-3x faster than sequential processing
- **Resource-Aware Runtime** - Auto-selects the appropriate Qwen3-ASR model for the current machine
## Acknowledgements
- [Qwen3-ASR](https://github.com/QwenLM/Qwen3-ASR) provides the official model family and multimodal/vLLM usage guidance
- [QwenASR](https://github.com/huanglizhuo/QwenASR) provides the CPU Rust backend vendored by this project
## Quick Deployment
### 1. Docker Deployment (Recommended)
```bash
# Copy and edit configuration
cp .env.example .env
# Edit .env to set API_KEY (optional)
# Compose defaults:
# /opt/dep/asr/models -> /app/models
# /opt/dep/asr/data -> /app/data
# /opt/dep/asr/data/logs, temp, tasks live under this data mount
# Optional: override any host mount root in .env
# export MODEL_STORAGE_DIR=/data/qwen3-asr-models
# export DATA_STORAGE_DIR=/data/qwen3-asr-data
# Start service (NVIDIA GPU version)
docker-compose up -d
# Or MetaX/MuXi GPU version
docker-compose -f docker-compose-metax.yml up -d
# Or Iluvatar/Tianshu GPU version
docker-compose -f docker-compose-iluvatar.yml up -d
# Or Moore Threads / MUSA GPU version
docker-compose -f docker-compose-mthreads.yml up -d
# Or CPU version
docker-compose -f docker-compose-cpu.yml up -d
# NVIDIA multi-GPU auto mode (one instance per visible GPU)
CUDA_VISIBLE_DEVICES=0,1,2,3 docker-compose up -d
# MetaX/MuXi multi-GPU auto mode
METAX_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-metax.yml up -d
# Iluvatar/Tianshu multi-GPU auto mode
ILUVATAR_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-iluvatar.yml up -d
# Moore Threads / MUSA multi-GPU auto mode
MTHREADS_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-mthreads.yml up -d
```
Service URLs:
- **API Endpoint**: `http://localhost:17003`
- **API Docs**: `http://localhost:17003/docs`
Optional built-in rate limit settings:
- `NGINX_RATE_LIMIT_RPS` (global requests/sec, `0` = disabled)
- `NGINX_RATE_LIMIT_BURST` (global burst, `0` = auto use RPS)
**docker run (alternative):**
```bash
# NVIDIA GPU version
docker run -d --name qwen3-asr \
--gpus all \
-p 17003:8000 \
-e ACCELERATOR=nvidia \
-e CUDA_VISIBLE_DEVICES=0,1,2,3 \
-e API_KEY=your_api_key \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:gpu-latest
# MetaX/MuXi GPU version
docker run -d --name qwen3-asr-metax \
--privileged \
--network=host \
--pid=host \
--ipc=host \
-v /dev:/dev \
-v /opt/mxdriver:/opt/mxdriver:ro \
-e ACCELERATOR=metax \
-e PORT=17003 \
-e METAX_VISIBLE_DEVICES=0 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:metax-latest
# CPU version
docker run -d --name qwen3-asr \
-p 17003:8000 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:cpu-latest
```
> **Note**: NVIDIA GPU images default to CUDA 13.0/cu130 with `torch 2.11.0` + `vllm 0.20.0`.
> Developers can rebuild `Dockerfile.gpu` for CUDA 12.6, CUDA 13.0, or another backend by overriding Docker build args.
> MetaX/MuXi images use `Dockerfile.metax` on top of an official MetaX vLLM image. In field deployments, use host networking plus privileged `/dev` and `/opt/mxdriver` mounts so both `mx-smi` and the MetaX PyTorch runtime can initialize devices.
> CPU images now support `qwen3-asr-0.6b` via the bundled QwenASR Rust backend. The default CPU image uses a portable Rust target; set `QWENASR_RUST_TARGET_CPU=native` only for self-built, host-specific images.
> On CUDA vLLM and CPU Rust, `word_timestamps=true` now triggers the forced aligner automatically.
> On macOS / Apple Silicon, Qwen3-ASR now runs through the Rust CPU backend.
> `start.py` now forces the vLLM multiprocessing method to `spawn` so startup does not hit CUDA re-initialization failures in forked subprocesses.
**Custom GPU backend builds:**
```bash
# Default GPU build: CUDA 13.0 / PyTorch cu130
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu .
# CUDA 12.6 build for older deployments
docker build -t qwen3-asr:gpu-cu126 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda12.6-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu126 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-12-6 \
--build-arg TORCH_CUDA_ARCH_LIST="8.0;8.6;8.9" \
.
# CUDA 13.0 build when your driver/toolchain requires it
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda13.0-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu130 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-13-0 \
--build-arg TORCH_CUDA_ARCH_LIST="12.0+PTX" \
.
# MetaX/MuXi build: fuse this project into an official MetaX vLLM image
./scripts/package_vendor_gpu_image.sh \
--vendor metax \
--base-image <official-metax-vllm-image> \
-v n260-3.7.0.38
# Iluvatar/Tianshu build: fuse this project into the official Iluvatar vLLM image
docker pull registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
./scripts/package_vendor_gpu_image.sh \
--vendor iluvatar \
--base-image registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5
# Moore Threads / MUSA build: fuse this project into the official MUSA vLLM image
docker pull registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
./scripts/package_vendor_gpu_image.sh \
--vendor mthreads \
--base-image registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519 \
-v s4000_4.3.5_d0519
```
For MetaX/MuXi offline delivery, see [docs/metax_offline_deployment.md](docs/metax_offline_deployment.md).
For Iluvatar/Tianshu offline delivery, see [docs/iluvatar_offline_deployment.md](docs/iluvatar_offline_deployment.md).
For Moore Threads / MUSA offline delivery, see [docs/mthreads_offline_deployment.md](docs/mthreads_offline_deployment.md).
**Offline Deployment**: You can now build a timestamped offline delivery folder that includes the image archive, compose file, env template, host-dir init script, and usage docs. The export script uses plain `docker build` + `docker save`, so it does not depend on `buildx`:
```bash
# 1. Build an offline delivery folder
./export_offline_bundle.sh --type gpu
# or
./export_offline_bundle.sh --type cpu
# or MetaX/MuXi GPU
./export_offline_bundle.sh \
--type metax \
--metax-base cr.metax-tech.com/public-ai-release/maca/vllm-metax:0.17.0-maca.ai3.5.3.307-torch2.8-py312-ubuntu22.04-amd64 \
--skip-models
# or Iluvatar/Tianshu GPU
./export_offline_bundle.sh --type iluvatar --iluvatar-base registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
# or Moore Threads / MUSA GPU
./export_offline_bundle.sh --type mthreads --mthreads-base registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
# or build both in one bundle
./export_offline_bundle.sh --type all
# 2. Prepare models separately, without deleting existing model files
./scripts/download-models.sh --models-dir /opt/dep/asr/models
# 3. Copy the generated folder to the offline server
scp -r build-file/<timestamp>-all user@server:/opt/dep/asr/
# 4. On the offline server
cd /opt/dep/asr/<timestamp>-all
./init_host_dirs.sh
gunzip -c qwen3-asr-gpu-<timestamp>-amd64.tar.gz | docker load
gunzip -c qwen3-asr-cpu-<timestamp>-amd64.tar.gz | docker load
# NVIDIA GPU
docker compose up -d
# or MetaX/MuXi GPU
# docker compose -f docker-compose-metax.yml up -d
# or Iluvatar/Tianshu GPU
# docker compose -f docker-compose-iluvatar.yml up -d
# or Moore Threads / MUSA GPU
# docker compose -f docker-compose-mthreads.yml up -d
# or CPU
# docker compose -f docker-compose-cpu.yml up -d
```
> Detailed deployment instructions: [Deployment Guide](./docs/deployment.md)
### Local Development
**System Requirements:**
- Python 3.10+
- CUDA 13.0+ for the default GPU image; CUDA 12.6 / 13.0 can be built with Docker args
- FFmpeg (audio format conversion)
**Installation:**
Runtime dependency locks now default to the GPU stack at the repo root, with CPU kept as a specialized environment:
| Mode | Command | Notes |
|------|---------|-------|
| NVIDIA GPU (default) | `uv sync` or `./scripts/sync_gpu_env.sh` | Syncs the root [pyproject.toml](/opt/qwen3-asr/pyproject.toml) and [uv.lock](/opt/qwen3-asr/uv.lock) into `.venv`, including CUDA 13.0/cu130 `torch 2.11.0` / `torchaudio 2.11.0` / `torchvision 0.26.0` / `vllm 0.20.0` |
| MetaX/MuXi GPU | `./scripts/sync_metax_env.sh` | Syncs common dependencies from [environments/metax/pyproject.toml](/opt/qwen3-asr/environments/metax/pyproject.toml); optional GPU-stack install uses the MetaX MACA PyPI index with `--no-deps` by default |
| Iluvatar/Tianshu GPU | `./scripts/sync_iluvatar_env.sh` | Syncs common dependencies from [environments/iluvatar/pyproject.toml](/opt/qwen3-asr/environments/iluvatar/pyproject.toml); GPU stack should come from the official Iluvatar vLLM image |
| Moore Threads / MUSA GPU | `./scripts/sync_mthreads_env.sh` | Syncs common dependencies from [environments/mthreads/pyproject.toml](/opt/qwen3-asr/environments/mthreads/pyproject.toml); GPU stack should come from the official Moore Threads MUSA vLLM image |
| CPU (specialized) | `./scripts/sync_cpu_env.sh` | Syncs the dedicated CPU lock in [environments/cpu/pyproject.toml](/opt/qwen3-asr/environments/cpu/pyproject.toml) into `.venv` |
| Auto | `./scripts/sync_accel_env.sh` | Chooses MetaX when `mx-smi` is present, Iluvatar when `ixsmi` is present, Moore Threads when `mthreads-gmi` is present, otherwise NVIDIA when `nvidia-smi` is present, otherwise CPU |
```bash
# Clone project
cd qwen3-asr
# Install dependencies (Linux/NVIDIA CUDA)
uv sync
# Start service
source .venv/bin/activate
python start.py
```
MetaX/MuXi local development:
```bash
./scripts/sync_metax_env.sh
source .venv/bin/activate
ACCELERATOR=metax python start.py
```
Local model storage defaults to `./models` under the project root:
```text
./models/
Qwen/
iic/
damo/
```
Override it when needed:
```bash
export MODELS_DIR=/data/qwen3-asr-models
export MODELSCOPE_CACHE=/data
export MODELSCOPE_PATH=$MODELS_DIR
```
macOS / Apple Silicon local development:
```bash
./scripts/sync_cpu_env.sh
source .venv/bin/activate
python start.py
```
## Runtime Defaults
Current runtime behavior on the mainline codebase:
- `ACCELERATOR=auto` resolves to `metax` when `mx-smi` reports devices, then `iluvatar` when `ixsmi` reports devices, then `mthreads` when `mthreads-gmi` reports devices, otherwise `nvidia` when NVIDIA CUDA is available, otherwise `cpu`
- `DEVICE=auto` resolves to the active accelerator device (`cuda:0` for NVIDIA/MetaX/Iluvatar GPU, otherwise `cpu`)
- `DEVICE=mps` is normalized to `cpu`
- `Linux + NVIDIA CUDA` uses official `vLLM`
- `Linux + MetaX/MuXi MACA` uses the MetaX-compatible PyTorch/vLLM stack
- `Linux + Iluvatar/Tianshu` uses the Iluvatar official vLLM image stack
- `Linux + CPU` uses vendored `QwenASR` Rust
- `macOS / Apple Silicon` also uses vendored `QwenASR` Rust
- macOS / Apple Silicon defaults to `qwen3-asr-0.6b`
- `qwen3-asr-1.7b` on macOS is only used when `QWEN3_ASR_MODEL=qwen3-asr-1.7b`
- `word_timestamps=true` works on the current offline CUDA and CPU Rust paths
- WebSocket streaming does not currently return word-level timestamps
- CAM++ speaker diarization remains required and still follows `DEVICE`; on CPU its main hotspot is speaker verification embedding
## API Endpoints
### OpenAI Compatible API
| Endpoint | Method | Function |
|----------|--------|----------|
| `/v1/audio/transcriptions` | POST | Audio transcription (OpenAI compatible) |
| `/v1/models` | GET | Offline model list |
**Request Parameters:**
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `file` | file | Preferred when provided | Audio/video file |
| `audio_address` | string | Optional | Audio/video URL (HTTP/HTTPS), `file://`, or server-local path. Ignored when `file` is also provided |
| `language` | string | Auto-detect | Language code (zh/en/ja) |
| `enable_speaker_diarization` | bool | `true` | Enable speaker diarization |
| `enable_speaker_identification` | bool | `true` | Match registered speaker database when diarization is enabled |
| `enable_text_cleanup` | bool | `true` | Enable text deduplication, boundary-overlap trimming, and filler cleanup |
| `word_timestamps` | bool | `false` | Return word-level timestamps when the backend supports them. Qwen CUDA vLLM and CPU Rust automatically use the forced aligner when enabled. |
| `hotwords` | string | - | Hotwords, format: `word1 weight1 word2 weight2` |
| `response_format` | string | `verbose_json` | Output format |
| `prompt` | string | - | Prompt text (reserved) |
| `temperature` | float | `0` | Sampling temperature (reserved) |
**Audio / Video Input Methods:**
- **File Upload**: Use `file` parameter to upload an audio file or a video container with an audio track
- **URL / Local Path**: Use `audio_address` parameter to provide an audio/video URL or server-local path, service will read it automatically
- **Precedence**: If both `file` and `audio_address` are provided, the service uses `file` and ignores `audio_address`
**Usage Examples:**
```python
# Using OpenAI SDK
from openai import OpenAI
client = OpenAI(base_url="http://localhost:8000/v1", api_key="your_api_key")
with open("audio.wav", "rb") as f:
transcript = client.audio.transcriptions.create(
file=f,
response_format="verbose_json" # Get segments and speaker info
)
print(transcript.text)
```
```bash
# Using curl
curl -X POST "http://localhost:8000/v1/audio/transcriptions" \
-H "Authorization: Bearer your_api_key" \
-F "file=@audio.wav" \
-F "model=qwen3-asr-0.6b" \
-F "response_format=verbose_json" \
-F "enable_speaker_diarization=true" \
-F "enable_speaker_identification=true" \
-F "enable_text_cleanup=true" \
-F "hotwords=Qwen 2.0 ModelScope 1.5"
```
**Supported Response Formats:** `json`, `text`, `srt`, `vtt`, `verbose_json`
### Alibaba Cloud Compatible API
| Endpoint | Method | Function |
|----------|--------|----------|
| `/stream/v1/asr` | POST | Speech recognition (long audio support) |
| `/stream/v1/asr/models` | GET | Declared model/capability entries |
| `/stream/v1/asr/health` | GET | Health check |
| `/ws/v1/asr` | WebSocket | Qwen3-ASR streaming |
| `/ws/v1/asr/qwen` | WebSocket | Qwen3-ASR streaming (explicit path) |
| `/ws/v1/asr/funasr` | WebSocket | Removed; returns a deprecation error and asks clients to switch to `/ws/v1/asr/qwen` |
**Request Parameters:**
| Parameter | Type | Default | Description |
|-----------|------|---------|-------------|
| `audio_address` | string | `https://media.cdn.vect.one/podcast_demo.mp4` (docs example) | Audio/video URL, `file://`, or server-local path (optional; ignored when body content is uploaded) |
| `sample_rate` | int | `16000` | Sample rate |
| `enable_speaker_diarization` | bool | `true` | Enable speaker diarization |
| `enable_speaker_identification` | bool | `true` | Match registered speaker database when diarization is enabled |
| `enable_text_cleanup` | bool | `true` | Enable text deduplication, boundary-overlap trimming, and filler cleanup |
| `word_timestamps` | bool | `false` | Return word-level timestamps when the backend supports them. Qwen CUDA vLLM and CPU Rust automatically use the forced aligner when enabled. |
| `vocabulary_id` | string | - | Hotwords (format: `word1 weight1 word2 weight2`) |
**Usage Examples:**
```bash
# Basic usage
curl -X POST "http://localhost:8000/stream/v1/asr" \
-H "Content-Type: application/octet-stream" \
--data-binary @audio.wav
# With parameters
curl -X POST "http://localhost:8000/stream/v1/asr?enable_speaker_diarization=true&enable_speaker_identification=true&enable_text_cleanup=true&vocabulary_id=Qwen%202.0%20ModelScope%201.5" \
-H "Content-Type: application/octet-stream" \
--data-binary @audio.wav
```
### Meeting Offline API
| Endpoint | Method | Function |
|----------|--------|----------|
| `/api/v1/asr/transcriptions` | POST | Create an offline meeting transcription task |
| `/api/v1/asr/transcriptions/{task_id}` | GET | Query task status and result |
`audio_address` is the required production input field for this endpoint.
```json
{
"audio_address": "https://example.com/media/meeting.mp4",
"config": {
"enable_speaker": true,
"match_speaker_registry": true,
"enable_text_cleanup": true,
"speaker_threshold": 0.6,
"word_timestamps": false,
"hotwords": [
{ "hotword": "Qwen", "weight": 2.0 },
{ "hotword": "ModelScope", "weight": 1.5 }
]
}
}
```
**Response Example:**
```json
{
"task_id": "xxx",
"status": 200,
"message": "SUCCESS",
"result": "Speaker1 content...\nSpeaker2 content...",
"duration": 60.5,
"processing_time": 1.234,
"segments": [
{
"text": "Today is a nice day.",
"start_time": 0.0,
"end_time": 2.5,
"speaker_id": "Speaker1",
"word_tokens": [
{"text": "Today", "start_time": 0.0, "end_time": 0.5},
{"text": "is", "start_time": 0.5, "end_time": 0.7},
{"text": "a nice day", "start_time": 0.7, "end_time": 1.5}
]
}
]
}
```
## Speaker Diarization
Multi-speaker automatic identification based on CAM++ model:
- **Enabled by Default** - `enable_speaker_diarization=true`
- **Automatic Detection** - No preset speaker count needed, model auto-detects
- **Speaker Labels** - Response includes `speaker_id` field (e.g., "Speaker1", "Speaker2")
- **Smart Merging** - Two-layer merge strategy to avoid isolated short segments:
- Layer 1: Accumulate merge same-speaker segments < 10 seconds
- Layer 2: Accumulate merge continuous segments up to 60 seconds
- **Subtitle Support** - SRT/VTT output includes speaker labels `[Speaker1] text content`
Disable speaker diarization:
```bash
# OpenAI API
-F "enable_speaker_diarization=false"
# Alibaba Cloud API
?enable_speaker_diarization=false
```
## Audio Processing
### Intelligent Segmentation Strategy
Automatic long audio segmentation:
1. **VAD Voice Detection** - Detect voice boundaries, filter silence
2. **Greedy Merge** - Accumulate voice segments, ensure each segment does not exceed `MAX_SEGMENT_SEC` (default 60s)
3. **Silence Split** - Force split when silence between voice segments exceeds 3 seconds
4. **Batch Inference** - Multi-segment parallel processing, 2-3x performance improvement in GPU mode
### WebSocket Streaming Limitations
**Qwen3-ASR Streaming** (using `/ws/v1/asr` or `/ws/v1/asr/qwen`):
- ✅ Multi-language real-time recognition
- ✅ CUDA vLLM and CPU Rust both support the current streaming path
- ❌ Word-level timestamps are not available in the current streaming path
### Qwen3 Runtime Matrix
| Runtime | Backend | Offline | WebSocket Streaming | Word Timestamps Offline | Word Timestamps Streaming | Maturity |
|---------|---------|---------|---------------------|-------------------------|---------------------------|----------|
| Linux + NVIDIA GPU | Official vLLM 0.20.0 | ✅ | ✅ | ✅ | ❌ | Production-oriented |
| CPU / macOS | QwenASR Rust | ✅ | ✅ | ✅ (forced aligner) | ❌ | Recommended local fallback |
## Offline-Capable Models
| Model ID | Name | Description | Features |
|----------|------|-------------|----------|
| `qwen3-asr-1.7b` | Qwen3-ASR 1.7B | High-performance multilingual ASR, 52 languages + dialects; CUDA uses vLLM | Offline/Realtime |
| `qwen3-asr-0.6b` | Qwen3-ASR 0.6B | Lightweight multilingual ASR; CUDA uses vLLM, CPU/macOS uses Rust backend | Offline/Realtime |
**Runtime selection:**
- **VRAM >= 32GB**: Select `qwen3-asr-1.7b`
- **VRAM < 32GB**: Select `qwen3-asr-0.6b`
- **No CUDA**: Select the vendored Rust-backed `qwen3-asr-0.6b`
- **macOS / Apple Silicon**: Always default to `qwen3-asr-0.6b`, regardless of memory size
- **Environment override**: Set `QWEN3_ASR_MODEL=qwen3-asr-1.7b` or `QWEN3_ASR_MODEL=qwen3-asr-0.6b` to bypass automatic selection
At startup the service checks the current runtime model plan and downloads missing models from ModelScope by default.
## Environment Variables
Recommended public settings:
| Variable | Default | Description |
|----------|---------|-------------|
| `API_KEY` | - | API authentication key (optional, unauthenticated if not set) |
| `LOG_LEVEL` | `INFO` | Log level (DEBUG/INFO/WARNING/ERROR) |
| `MAX_AUDIO_SIZE` | `2048` | Max audio file size (MB, supports units like 2GB) |
| `ASR_BATCH_SIZE` | `4` | ASR batch size for long-audio segment processing |
| `MAX_SEGMENT_SEC` | `60` | Max audio segment duration (seconds) |
| `ASR_ENABLE_NEARFIELD_FILTER` | `true` | Enable far-field sound filtering |
| `QWEN3_ASR_MODEL` | auto | Force `qwen3-asr-1.7b` or `qwen3-asr-0.6b` instead of VRAM-based selection |
| `QWEN_GPU_MEMORY_UTILIZATION` | `0.9` | Upper bound for vLLM GPU memory reservation; lower it on shared GPUs, raise it when KV cache is too small |
| `QWEN_VLLM_ENFORCE_EAGER` | `true` | Force vLLM eager execution for compatibility; set `false` to allow CUDA Graph optimization on supported NVIDIA deployments |
Far-field filter notes:
- `ASR_NEARFIELD_RMS_THRESHOLD=0.01` is the current default and recommended starting point
- raise it in noisy rooms to filter more background speech
- lower it in quiet rooms if soft speech is being dropped
- use `LOG_LEVEL=DEBUG` temporarily when you need to inspect filter behavior
Advanced backend-specific settings:
| Variable | Default | Description |
|----------|---------|-------------|
| `QWEN_RUST_CPU_WORKERS` | `4` | CPU Rust backend worker count (Rust ASR / forced align default to 4 runtimes) |
| `QWENASR_LIBRARY_PATH` | auto-detect | Override vendored Rust dylib/so path |
## Resource Requirements
**Minimum (CPU):**
- CPU: 4 cores
- Memory: 16GB
- Disk: 20GB
**Recommended (GPU):**
- CPU: 4 cores
- Memory: 16GB
- GPU: NVIDIA GPU (16GB+ VRAM)
- Disk: 20GB
## API Documentation
After starting the service:
- Swagger UI: `http://localhost:8000/docs`
- ReDoc: `http://localhost:8000/redoc`
## Links
- **Deployment Guide**: [Detailed Docs](./docs/deployment.md)
- **Qwen3-ASR**: [Qwen3-ASR GitHub](https://github.com/QwenLM/Qwen3-ASR)
- **FunASR**: [FunASR GitHub](https://github.com/alibaba-damo-academy/FunASR)
- **Chinese README**: [中文文档](./docs/README_zh.md)
## License
This project uses the MIT License - see [LICENSE](LICENSE) file for details.
## Star History
[![Star History Chart](https://api.star-history.com/svg?repos=Quantatirsk/qwen3-asr&type=Date)](https://star-history.com/#Quantatirsk/qwen3-asr&Date)
## Contributing
Issues and Pull Requests are welcome to improve the project!

View File

@ -0,0 +1,310 @@
# 实时 ASR 当前问题排查记录
更新时间:2026-09-07
排查范围:实时 ASR、原生流式 partial、说话人识别、声纹姓名匹配、WebSocket 输出和前端展示。
## 1. 当前结论
目前发现的问题分布在三个层次:
```text
识别层 原生 partial 默认关闭,非原生路径会反复重识别窗口
说话人层 短片段无条件继承上一位实名,并可能把复制的 embedding 写回记录池
输出层 内部物理切段直接作为前端展示单元,导致同一说话人被拆成多行
```
其中,“不同的人进入已经确定的说话人气泡”最明确的根因是说话人层的短段快路径:小于 1.6 秒的片段在提取新特征之前直接继承上一位已命名说话人。WebSocket 本身负责传递这些结果,但错误身份是在上游状态机中产生的。
“同一说话人被切成几行”主要是输出层问题。silence 和 max duration 可以结束一次内部识别片段,但不应直接决定前端展示换行。
## 2. 当前实时 WebSocket 链路
入口位于 [app/api/v1/websocket_asr.py](app/api/v1/websocket_asr.py)。`/ws/v1/asr` 和 `/ws/v1/asr/qwen` 最终都进入 [Qwen3ASRService.handle_connection](app/services/qwen3_websocket_asr.py#L2538)。
主流程如下:
```text
客户端 start
↓
初始化 ConnectionContext 和 streaming state
↓
接收二进制音频
↓
RMS/peak 判断是否有声音,维护 pre-roll 和当前 turn
↓
partial 解码并向前端发送 sentence_type=0
↓
silence / max_duration / sentence_limit / stop 触发内部提交
↓
最终重识别,创建 confirmed segment
↓
异步 speaker worker 处理说话人
↓
通过相同 sentence_id 发送 speaker update
↓
stop 时发送最终 sentences 和 end
```
图谱分析显示 `handle_connection` 是当前实时链路的关键汇聚点;说话人解析又连接到 `RealtimeSpeakerClusterer.resolve_segment_speaker`、声纹注册服务和历史记录重聚类。因此直接修改主服务的影响范围较大,本次验证优先放在独立 demo 中。
## 3. 问题一:内部切段直接变成前端展示行
### 3.1 当前切段触发器
| 触发器 | 当前条件 | 当前作用 | 对展示的实际影响 |
|---|---|---|---|
| `silence` | 连续静音默认 `800ms` | 主力结束当前 turn | 直接产生一个 confirmed segment |
| `max_duration` | 缓冲默认最多 `12s` | 防止单段无限增长 | 长讲话被硬切成多个 segment |
| `sentence_limit` | 段内达到默认 `8` 句 | 兜底限制 | 可能提前提交 |
| `complete_sentence` | 标点结尾、时长达到阈值且字数足够 | 长独讲兜底 | 正常短句通常到不了该条件 |
| `final` | 客户端发送 `stop` | 提交最后一段 | 输出最终结果 |
相关逻辑位于 [qwen3_websocket_asr.py](app/services/qwen3_websocket_asr.py#L2538)、[_should_commit_complete_sentence](app/services/qwen3_websocket_asr.py#L1100) 和 [_commit_retranscribe_turn](app/services/qwen3_websocket_asr.py#L2333)。
### 3.2 当前输出方式
每次内部提交都会:
1. 创建一条 `confirmed_segments` 记录;
2. 使用该记录的 `index` 作为 `sentence_id`;
3. 立即发送 `sentence_type=1` 的 `sentences` 事件;
4. 把所有记录用换行连接成 `full_text`。
代码中 `full_text` 使用 `"\\n".join(...)`,位置在 [_commit_retranscribe_turn](app/services/qwen3_websocket_asr.py#L2464) 和 [_stop](app/services/qwen3_websocket_asr.py#L2916)。因此这里的 `sentence_id` 实际上是“物理切段编号”,不一定是语义完整句子编号。
项目协议文档明确区分了 `sentence_type=0` 的 partial 和 `sentence_type=1` 的 final,但没有把“内部 segment”和“前端 display block”分开,见 [realtime_meeting_websocket.md](docs/realtime_meeting_websocket.md#L162)。
前端可以根据相同 `sentence_id` 更新原记录,但不会把不同 `sentence_id` 且说话人相同的记录合并。因此停顿、12 秒硬切和前端换行目前形成了直接关系。
另外,图谱显示 [_should_force_stable_segment](app/services/qwen3_websocket_asr.py#L2051) 当前没有调用方,它不能实际改变实时切段;`complete_sentence` 主要是长段兜底。
## 4. 问题二:原生流式 partial 与输出合并不是同一层
当前配置中 `REALTIME_STREAM_CHUNK_SEC` 默认约为 `1.2s`,`REALTIME_PARTIAL_WINDOW_SEC` 默认约为 `8s`,原生 partial 开关默认关闭,相关默认值在 [app/core/config.py](app/core/config.py#L103)。
非原生路径会维护窗口并调用 `_transcribe_audio_text` 重识别;原生路径则通过:
```text
create_stream → push_stream/feed_stream → finish_stream
```
对应 [qwen3_engine.py](app/services/asr/qwen3_engine.py#L570) 和 [qwen3_websocket_asr.py](app/services/qwen3_websocket_asr.py#L1162)。
原生 partial 影响的是:
- partial 首次出现的延迟;
- 每次增量识别的计算量;
- 流式状态维护方式;
- 文本回滚和未固定 token 的处理。
它不应决定:
- 是否因为停顿产生前端新行;
- 哪些物理 segment 合并成一个说话人气泡;
- 是否把某个 speaker 姓名写入展示结果。
所以两个优化需要同时做,但必须保持两个独立状态:ASR streaming state 和 display aggregation state。
## 5. 问题三:短音频无条件继承实名
### 5.1 已确认的代码路径
在 [realtime_speaker_clusterer.py](app/services/realtime_speaker_clusterer.py#L27) 中,`_FAST_ATTACH_MAX_SEC = 1.6`。
`resolve_segment_speaker` 的入口逻辑是:
```text
duration_sec < 1.6s
且 speaker_records 非空
且上一条记录存在实名身份
↓
直接复制上一条记录的 speaker_id / speaker_name / user_id / registry_speaker_id
speaker_confidence = 0.0
speaker_strategy = "short_attach"
```
该分支在 [realtime_speaker_clusterer.py](app/services/realtime_speaker_clusterer.py#L505),早于 `extract_chunk_embeddings`。因此短段没有经过新的声纹特征验证。
### 5.2 为什么实名会直接通过稳定性判断
[_is_stable_speaker_info](app/services/qwen3_websocket_asr.py#L613) 首先检查 `registry_speaker_id` 或 `user_id`。短段继承时这两个字段被完整复制,所以即使 `speaker_confidence=0.0`,仍会被判断为 stable。
随后 [_resolve_segment_speaker](app/services/qwen3_websocket_asr.py#L2156) 会把该身份写入最终 segment,并由 [_emit_speaker_update](app/services/qwen3_websocket_asr.py#L498) 通过相同 `sentence_id` 推送给客户端。
这解释了为什么匿名 `SpeakerNN` 的继承可能被拦截,而带 `user_id` 或注册声纹 ID 的实名继承可以直接进入前端。
### 5.3 为什么短段是高风险场景
“嗯”“对”“好的”“可以”等短插话通常不到 1.6 秒,恰好是多人会议中最容易发生换人的场景。当前快路径把上一位实名当作新片段身份,产生的结果就是:新说话人的短句进入上一位实名气泡。
## 6. 问题四:特征提取失败分支与实际路径
代码中确实存在 `embedding_attach` 分支:特征提取后如果 `current_chunks` 为空,则尝试继承上一位实名,位置在 [realtime_speaker_clusterer.py](app/services/realtime_speaker_clusterer.py#L521)。
但当前 `extract_chunk_embeddings` 在没有 chunk 时会回退为整段单 chunk,因此正常情况下很难返回空列表;如果模型真正抛异常,异常会向上传递,最终由 [_resolve_and_emit_segment_speaker](app/services/qwen3_websocket_asr.py#L2220) 捕获,当前片段保持 pending。
准确结论是:
- 模型抛异常:当前片段通常保持 `speaker_id=-1`,不会走继承;
- chunk 为空:代码意图是继承实名,但该分支近乎不可达;
- 模型不抛异常但产生垃圾 embedding:仍需单独检查零向量、NaN 和相似度边界;
- 当前最确定、最直接的错误来源是 `<1.6s` 的 `short_attach`。
## 7. 问题五:继承 embedding 造成污染链
短段继承返回时还会复制上一条记录的 `_embedding`、`_chunk_embeddings` 和 `_chunks`。下游 [_record_segment_speaker](app/services/qwen3_websocket_asr.py#L644) 只要发现 `_embedding` 非空,就会把它作为正常 speaker record 保存。
污染链如下:
```text
上一位实名离场
↓
新人的短段直接复制实名和 embedding
↓
复制结果成为新的 last_record
↓
后续短段继续继承这条记录
↓
历史匹配池重复出现同一个 embedding
↓
相似度、聚类中心和重聚类结果被污染
```
影响包括:
1. 继承链持续延长,错误实名被不断续写;
2. 历史 embedding 被重复计数,匹配置信度可能虚高;
3. `cluster_records_with_ranges` 会把复制的 chunk 作为真实样本参与聚类;
4. 后续重聚类可能把错误身份回写到更多 segment。
## 8. 问题六:speaker update 是异步的
实时最终段先写入 `confirmed_segments`,speaker worker 再异步处理,处理完成后通过相同 `sentence_id` 发送更新,相关位置是 [_speaker_worker_loop](app/services/qwen3_websocket_asr.py#L442) 和 [_emit_speaker_update](app/services/qwen3_websocket_asr.py#L498)。
当前协议允许首次 final 没有姓名,后续再补姓名,这一点在 [realtime_meeting_websocket.md](docs/realtime_meeting_websocket.md#L240) 有说明。
这里有两个风险:
- 前端如果把每次事件当成追加消息,会出现重复行;
- 上游异步继承结果如果带着错误实名,前端的幂等更新会把错误身份稳定显示出来。
因此上层 WebSocket 需要维护 segment 状态表,按 `sentence_id` 覆盖更新,然后根据完整状态重新生成 display block 快照。
## 9. 需要同时实施的两层优化
### 9.1 识别层:原生流式 partial
目标是让同一轮语音持续使用一个 streaming state,通过增量音频推进识别,减少窗口重复重识别。
demo 应透传并记录:
- `enable_native_partial_stream`;
- partial 产生时间;
- 每次 partial 的文本长度和修订次数;
- 首次 partial 延迟;
- final 延迟;
- 原服务返回的 chunk 或 segment 时间范围。
### 9.2 说话人与输出层
目标是把“身份确认”和“展示合并”分开:
1. `<1.6s` 片段不继承上一位实名;
2. 没有新鲜声纹证据时保持 pending;
3. 特征提取异常时保持 pending;
4. 不复制上一条记录的 embedding;
5. 只有达到确认阈值的独立特征才能更新身份缓存;
6. 同一 `sentence_id` 的更新覆盖原片段;
7. 相邻且身份可信度一致的物理 segment 才合并为 display block;
8. A→B→A 保留时间顺序,不把非相邻发言重新拼接到一起;
9. 前端展示使用 display block,入库和诊断仍保留 raw segment。
推荐的状态关系是:
```text
raw segment
├─ ASR text state:partial / final
├─ speaker evidence:pending / fresh / confirmed
├─ speaker identity:cluster / registry / user
└─ display block:按时间和可信身份重新生成
```
## 10. 服务器接口边界
已验证服务器 `10.100.53.199:8000` 可访问,原实时 WebSocket 地址为:
```text
ws://10.100.53.199:8000/ws/v1/asr/qwen
```
当前公开接口包括:
- `/ws/v1/asr/qwen`:原实时 WebSocket,已经包含主服务内部的 ASR、切段和 speaker 逻辑;
- `/v1/audio/transcriptions`:整段音频转写;
- `/api/v1/speakers/identify`:通过文件来源识别注册说话人;
- `/api/v1/speakers`:声纹注册和人员管理。
当前公开接口没有返回实时聚类所需的原始 embedding,也没有把 [Qwen3ASREngine](app/services/asr/qwen3_engine.py#L570) 的 `create_stream`、`feed_stream`、`finish_stream` 暴露为独立远程模型 RPC。
所以存在一个边界:
- 上层 demo WebSocket 可以独立重写事件顺序、切段提交策略、pending 保护和展示合并;
- 如果要在 demo 中完整重建原项目的声纹聚类,服务器还需要提供 embedding 接口或模型 RPC;
- 只依赖当前 `/ws/v1/asr/qwen` 返回字段,无法重新计算已被原服务错误归类的长段身份。
## 11. 当前独立 demo 状态
独立验证项目位于 [realtime_asr_optimization_demo](realtime_asr_optimization_demo)。
当前结构:
```text
浏览器
↓ ws://127.0.0.1:8082/ws
demo 上层 WebSocket
↓ model_service.py 远程模型服务适配层
↓ ws://10.100.53.199:8000/ws/v1/asr/qwen
已部署模型服务
```
当前 demo 已具备:
- 原生 partial 参数透传;
- 同一 `sentence_id` 覆盖更新;
- raw event 与 display block 分离;
- 同一说话人的相邻片段合并;
- `<1.6s` 实名结果降级为 pending;
- 已命名身份不接受只有弱 cluster id 的片段加入;
- embedding 字段不进入 demo 状态池。
- 模型服务与 WebSocket 解耦;demo 不启动或加载模型;
- 模型服务地址可通过页面或 `--model-service-url` 配置。
对应实现见 [server.py](realtime_asr_optimization_demo/server.py)、[model_service.py](realtime_asr_optimization_demo/model_service.py)、[speaker_assembler.py](realtime_asr_optimization_demo/speaker_assembler.py) 和 [static/app.js](realtime_asr_optimization_demo/static/app.js)。
这个 demo 当前主要验证上层协议、状态覆盖和展示保护。后续把模型服务部署到新服务器时,只需替换模型服务地址;若新模型服务协议不同,则替换 `model_service.py`,不改变 WebSocket 编排层。若要验证“demo 自己完成声纹特征提取、聚类和按需姓名匹配”,仍需服务器提供远程 embedding/model RPC。
## 12. 建议验证用例
| 用例 | 关注结果 |
|---|---|
| 单人连续讲话,中间停顿 800ms 以上 | 前端仍属于同一个 display block |
| A 讲话后,B 说“嗯/好的” | B 的短段保持 pending,不进入 A 的实名块 |
| A→B→A | 保持三个时间顺序块 |
| 同一 `sentence_id` 先 partial 后 final | 文本覆盖,不重复追加 |
| final 先返回,speaker update 后返回 | 原 block 更新并重新归并 |
| 声纹模型异常 | 片段保持 pending,不继承上一位实名 |
| 连续多个短插话 | 不复制上一条 embedding,不形成继承链 |
| 关闭 native partial | 只影响识别延迟和资源,不改变展示合并规则 |
| 关闭 display merge | 可以看到原始物理 segment,用于对照 |
## 13. 当前未确认项
以下问题不能仅靠现有公开 WebSocket 字段确认:
- 原服务是否存在未写入 OpenAPI 的内部 embedding/model RPC;
- 实际使用的 realtime speaker 模型是否会输出零向量或 NaN;
- 不同真实说话人被分配相同稳定 cluster id 的比例;
- speaker worker 完成时间与 `end` 事件之间是否存在竞态;
- 远程 Docker 是否还映射了独立的模型后端端口。
这些项目需要通过服务器日志、embedding 接口或带原始音频的端到端录音继续验证。

6
app/__init__.py 100644
View File

@ -0,0 +1,6 @@
# -*- coding: utf-8 -*-
"""Qwen3-ASR application package."""
__version__ = "1.0.1"
__author__ = "Nexa Team"
__description__ = "Qwen3-ASR speech recognition API service"

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
API路由模块
包含所有API端点的路由定义
"""

View File

@ -0,0 +1,22 @@
# -*- coding: utf-8 -*-
"""API v1版本路由"""
from fastapi import APIRouter
from .asr import router as asr_router
from .websocket_asr import router as websocket_asr_router
from .openai_compatible import router as openai_router
from .meeting import router as meeting_router
api_router = APIRouter()
# 原有 API (阿里云兼容)
api_router.include_router(asr_router)
# WebSocket ASR 端点(包含阿里云协议和 Qwen3 流式协议)
api_router.include_router(websocket_asr_router)
# OpenAI 兼容 API
api_router.include_router(openai_router)
# 独立会议离线/声纹管理 API(Model-Test-New 兼容形状)
api_router.include_router(meeting_router)

481
app/api/v1/asr.py 100644
View File

@ -0,0 +1,481 @@
# -*- coding: utf-8 -*-
"""
ASR API路由
"""
from fastapi import (
APIRouter,
Request,
HTTPException,
Depends
)
from fastapi.responses import JSONResponse
from typing import Annotated, Optional
import time
import logging
from ...core.config import settings
from ...core.exceptions import (
AuthenticationException,
InvalidParameterException,
InvalidMessageException,
UnsupportedSampleRateException,
DefaultServerErrorException,
get_http_status_code,
)
from ...core.security import validate_token
from ...models.common import SampleRate
from ...models.asr import (
ASRResponse,
ASRHealthCheckResponse,
ASRModelsResponse,
ASRSuccessResponse,
ASRErrorResponse,
ASRQueryParams,
)
from ...utils.common import generate_task_id
from ...services.asr.manager import get_model_manager
from ...services.asr.model_selection import validate_offline_model_id
from ...services.asr.runtime import get_runtime_router
from ...services.asr.audio_validation import validate_sample_rate
from ...services.asr.offline_transcription_service import (
OfflineTranscriptionOptions,
PreparedAudio,
get_offline_transcription_service,
)
# 配置日志
logger = logging.getLogger(__name__)
# 创建路由器
router = APIRouter(prefix="/stream/v1", tags=["ASR"])
def _build_asr_openapi_parameters() -> list[dict]:
parameters: list[dict] = [
{
"name": "model",
"in": "query",
"required": False,
"schema": {
"type": "string",
"maxLength": 128,
"example": "qwen3-asr-0.6b",
},
"description": "可选。离线 ASR 模型 ID;不传则使用服务当前默认模型",
},
{
"name": "audio_address",
"in": "query",
"required": False,
"schema": {
"type": "string",
"maxLength": 512,
"example": "https://media.cdn.vect.one/podcast_demo.mp4",
},
"description": "音频/视频文件地址,支持 HTTP/HTTPS URL、file:// 或服务端本地路径。仅当请求体为空时使用;若同时上传请求体,服务会忽略此参数",
},
{
"name": "sample_rate",
"in": "query",
"required": False,
"schema": {
"type": "integer",
"enum": [8000, 16000, 22050, 24000, 32000, 44100, 48000],
"default": 16000,
"example": 16000,
},
"description": "音频采样率(Hz)。音频会在服务端自动转换,通常保持默认值即可",
},
{
"name": "enable_speaker_diarization",
"in": "query",
"required": False,
"schema": {
"type": "boolean",
"default": True,
"example": True,
},
"description": "是否启用说话人分离。启用后响应会包含 speaker_id 字段",
},
{
"name": "enable_speaker_identification",
"in": "query",
"required": False,
"schema": {
"type": "boolean",
"default": True,
"example": True,
},
"description": "是否匹配已注册声纹库。仅在 enable_speaker_diarization=true 时生效,命中后响应会包含 speaker_name/user_id",
},
{
"name": "enable_text_cleanup",
"in": "query",
"required": False,
"schema": {
"type": "boolean",
"default": True,
"example": True,
},
"description": "是否启用文本去重、跨段重叠裁剪和口头语清理",
},
{
"name": "vocabulary_id",
"in": "query",
"required": False,
"schema": {
"type": "string",
"maxLength": 512,
"example": "阿里巴巴 20 腾讯 15",
},
"description": "热词字符串,格式:`热词1 权重1 热词2 权重2`。权重范围 1-100,建议 10-30。可提升特定词汇的识别准确率",
},
{
"name": "X-NLS-Token",
"in": "header",
"required": False,
"schema": {
"type": "string",
"minLength": 1,
"maxLength": 256,
"example": "",
},
"description": "访问令牌,用于身份认证。未配置 API_KEY 环境变量时可忽略",
},
]
if settings.ASR_ENABLE_WORD_TIMESTAMPS:
parameters.insert(
5,
{
"name": "word_timestamps",
"in": "query",
"required": False,
"schema": {
"type": "boolean",
"default": False,
"example": False,
},
"description": "是否返回字词级时间戳(默认关闭;启用时会自动调用 forced aligner)",
},
)
return parameters
async def get_asr_params(request: Request) -> ASRQueryParams:
"""从请求中提取并验证ASR参数"""
# 从URL查询参数中获取
query_params = dict(request.query_params)
# 使用统一的验证器验证参数
try:
# 验证采样率(转换为整数)
if "sample_rate" in query_params and query_params["sample_rate"]:
try:
sample_rate = int(query_params["sample_rate"]) # type: ignore
validated_rate = validate_sample_rate(sample_rate)
query_params["sample_rate"] = str(validated_rate) # type: ignore
except ValueError:
raise InvalidParameterException(
f"采样率必须是整数,收到: {query_params['sample_rate']}"
)
# 创建ASRQueryParams实例,Pydantic会自动验证和设置默认值
return ASRQueryParams.model_validate(query_params)
except InvalidParameterException:
raise
except Exception as e:
raise InvalidParameterException(f"请求参数错误: {str(e)}")
@router.post(
"/asr",
response_model=ASRResponse,
responses={
200: {
"description": "识别成功",
"model": ASRSuccessResponse,
},
400: {
"description": "请求参数错误",
"model": ASRErrorResponse,
},
401: {"description": "认证失败", "model": ASRErrorResponse},
500: {"description": "服务器内部错误", "model": ASRErrorResponse},
},
summary="语音识别(支持长音频)",
description="""
将音频文件转写为文本,兼容阿里云语音识别 RESTful API。
## 功能特性
- 支持多种音频格式与常见含音轨视频容器:WAV, MP3, M4A, FLAC, OGG, AAC, AMR, PCM, WEBM, MP4, MOV, MKV, AVI 等
- 自动音频格式检测和转换
- 支持长音频自动分段识别(返回带时间戳的分段结果)
- 最大文件大小:{settings.MAX_AUDIO_SIZE // (1024 * 1024)}MB(可通过环境变量 MAX_AUDIO_SIZE 配置)
## 音频输入方式
1. **请求体上传**:将音频/视频二进制数据作为请求体发送
2. **URL/本地路径读取**:通过 `audio_address` 参数指定音频/视频文件 URL(HTTP/HTTPS)或服务端本地路径
如果请求体和 `audio_address` 同时存在,服务会优先使用请求体,并忽略 `audio_address`。
## 注意事项
- 默认使用服务当前启用的 Qwen3-ASR 模型;也可通过可选 `model` 参数指定当前可用离线模型
- `vocabulary_id` 参数用于传递热词,格式:`热词1 权重1 热词2 权重2`(如:`阿里巴巴 20 腾讯 15`)
- `enable_speaker_identification` 仅在 `enable_speaker_diarization=true` 时生效,用于匹配已注册声纹库
- `enable_text_cleanup` 控制识别后的文本去重、跨段重叠裁剪和口头语清理
- 音频会自动转换为 16kHz 采样率进行识别
""",
openapi_extra={
"parameters": _build_asr_openapi_parameters(),
"requestBody": {
"description": "音频/视频文件二进制数据。支持格式:WAV, MP3, M4A, FLAC, OGG, AAC, AMR, PCM, WEBM, MP4, MOV, MKV, AVI 等。若同时提供 audio_address,服务会优先使用这里上传的内容",
"content": {
"application/octet-stream": {
"schema": {"type": "string", "format": "binary"}
}
},
"required": False,
},
},
)
async def asr_transcribe(
request: Request, params: Annotated[ASRQueryParams, Depends(get_asr_params)]
) -> JSONResponse:
"""语音识别API端点"""
task_id = generate_task_id()
prepared_audio: Optional[PreparedAudio] = None
# 性能计时
request_start_time = time.time()
# 记录请求开始(此时文件已上传完成)
content_length = request.headers.get("content-length", "unknown")
logger.info(f"[{task_id}] 收到ASR请求, content_length={content_length}")
transcription_service = get_offline_transcription_service()
try:
# 验证请求头部(鉴权)
result, content = validate_token(request, task_id)
if not result:
raise AuthenticationException(content, task_id)
model_id = validate_offline_model_id(params.model)
# 使用音频服务处理音频
target_sample_rate = int(params.sample_rate) if params.sample_rate else 16000
prepared_audio = await transcription_service.prepare_from_request(
request=request,
audio_address=params.audio_address,
task_id=task_id,
sample_rate=target_sample_rate,
)
logger.info(f"[{task_id}] 开始调用 transcribe_long_audio (enable_speaker_diarization={params.enable_speaker_diarization})...")
asr_result = await transcription_service.transcribe(
prepared_audio,
OfflineTranscriptionOptions(
model_id=model_id,
sample_rate=int(params.sample_rate or SampleRate.RATE_16000),
hotwords=params.vocabulary_id or "",
enable_speaker_diarization=params.enable_speaker_diarization is not False,
enable_speaker_identification=(
params.enable_speaker_diarization is not False
and params.enable_speaker_identification is not False
),
enable_text_cleanup=params.enable_text_cleanup is not False,
word_timestamps=(
settings.ASR_ENABLE_WORD_TIMESTAMPS
and params.word_timestamps is True
),
task_id=task_id,
),
)
logger.info(f"[{task_id}] 识别完成,共 {len(asr_result.segments)} 个分段,总字符: {len(asr_result.text)}")
# 构建分段结果(始终返回 segments,短音频也是 1 个 segment)
segments_data = []
for seg in asr_result.segments:
seg_dict = {
"text": seg.text,
"start_time": round(seg.start_time, 2),
"end_time": round(seg.end_time, 2),
}
if seg.speaker_id:
seg_dict["speaker_id"] = seg.speaker_id
if seg.speaker_name:
seg_dict["speaker_name"] = seg.speaker_name
if seg.user_id:
seg_dict["user_id"] = seg.user_id
# 添加字词级时间戳(如果存在)
if seg.word_tokens:
seg_dict["word_tokens"] = [
{
"text": wt.text,
"start_time": round(wt.start_time, 3),
"end_time": round(wt.end_time, 3),
}
for wt in seg.word_tokens
]
segments_data.append(seg_dict)
# 计算请求处理时间
request_duration = time.time() - request_start_time
# 返回成功响应(统一数据结构)
response_data = {
"task_id": task_id,
"result": asr_result.text,
"status": 200,
"message": "SUCCESS",
"segments": segments_data,
"duration": round(asr_result.duration, 2),
"processing_time": round(request_duration, 3),
}
return JSONResponse(content=response_data, headers={"task_id": task_id})
except (
AuthenticationException,
InvalidParameterException,
InvalidMessageException,
UnsupportedSampleRateException,
DefaultServerErrorException,
) as e:
e.task_id = task_id
logger.error(f"[{task_id}] ASR异常: {e.message}")
# 使用标准错误格式
response_data = e.to_dict()
return JSONResponse(
content=response_data,
headers={"task_id": task_id},
status_code=get_http_status_code(e.status_code),
)
except Exception as e:
logger.error(f"[{task_id}] 未知异常: {str(e)}")
# 使用标准错误格式
from ...core.exceptions import create_error_response
response_data = create_error_response(
error_code="DEFAULT_SERVER_ERROR",
message=f"内部服务错误: {str(e)}",
task_id=task_id,
)
return JSONResponse(content=response_data, headers={"task_id": task_id})
finally:
transcription_service.cleanup(prepared_audio)
@router.get(
"/asr/health",
response_model=ASRHealthCheckResponse,
summary="ASR 服务健康检查",
description="""
检查语音识别服务的运行状态和资源使用情况。
## 返回信息
- **status**: 服务状态(healthy/unhealthy/error)
- **model_loaded**: 默认模型是否已加载
- **device**: 当前推理设备(cuda:0/cpu)
- **loaded_models**: 已加载的模型列表
- **memory_usage**: GPU 显存使用情况(仅 GPU 模式)
""",
)
async def health_check(request: Request):
"""ASR服务健康检查端点"""
# 鉴权
result, content = validate_token(request)
if not result:
raise AuthenticationException(content, "health_check")
try:
# 尝试获取默认模型的引擎
try:
runtime_router = get_runtime_router()
default_model = runtime_router.resolve_model_id(None)
async with await runtime_router.acquire_engine(default_model) as engine:
model_loaded = True
device = engine.device
except Exception:
model_loaded = False
device = "unknown"
runtime_router = get_runtime_router()
memory_info = runtime_router.get_memory_usage()
loaded_models = runtime_router.get_loaded_model_ids()
accelerator_info = memory_info.get("accelerator")
return {
"status": "healthy" if model_loaded else "unhealthy",
"model_loaded": model_loaded,
"device": device,
"version": settings.APP_VERSION,
"message": (
"ASR service is running normally"
if model_loaded
else "ASR model not loaded"
),
"loaded_models": loaded_models,
"memory_usage": memory_info.get("gpu_memory"),
"accelerator": accelerator_info,
}
except Exception as e:
return {
"status": "error",
"model_loaded": False,
"device": "unknown",
"version": settings.APP_VERSION,
"message": str(e),
"accelerator": None,
}
@router.get(
"/asr/models",
response_model=ASRModelsResponse,
summary="获取声明条目列表",
description="""
返回系统声明的离线模型与 realtime capability 信息。
## 条目说明
| ID | 类型 | 说明 |
|----|------|------|
| qwen3-asr-1.7b | model | 离线/实时共用的 Qwen3-ASR 模型条目 |
| qwen3-asr-0.6b | model | 轻量版 Qwen3-ASR 模型条目 |
## 返回信息
- **declared_entries**: 声明的模型与 capability 列表
- **declared_count**: 声明项总数
- **runtime**: 运行时加载状态
""",
)
async def list_models(request: Request):
"""获取声明条目列表端点"""
# 鉴权
result, content = validate_token(request)
if not result:
raise AuthenticationException(content, "list_models")
try:
model_manager = get_model_manager()
runtime_router = get_runtime_router()
loaded_model_ids = runtime_router.get_loaded_model_ids()
entries = model_manager.list_declared_entries()
return {
"declared_entries": entries,
"declared_count": len(entries),
"runtime": {
"loaded_model_ids": loaded_model_ids,
"loaded_count": len(loaded_model_ids),
"default_offline_model_id": runtime_router.resolve_model_id(None),
},
}
except Exception as e:
logger.error(f"获取模型列表时发生错误: {str(e)}")
raise HTTPException(status_code=500, detail=f"获取模型列表失败: {str(e)}")

View File

@ -0,0 +1,855 @@
# -*- coding: utf-8 -*-
"""Independent meeting-style offline API compatible with Model-Test-New shapes."""
from __future__ import annotations
import asyncio
import math
import time
import uuid
from pathlib import Path
from typing import Any, Optional
from urllib.parse import unquote, urlparse
import requests
from fastapi import APIRouter, BackgroundTasks, Body, File, HTTPException, UploadFile
from pydantic import BaseModel, ConfigDict, Field
from app.core.config import settings
from app.core.database import pg_speaker_db
from app.core.exceptions import InvalidParameterException
from app.core.task_store import create_task_record, get_task_record, update_task_record
from app.models.common import SampleRate
from app.services.asr.model_selection import validate_offline_model_id
from app.services.asr.offline_transcription_service import (
OfflineTranscriptionOptions,
PreparedAudio,
get_offline_transcription_service,
)
from app.services.speaker_registry import get_speaker_registry_service
router = APIRouter(prefix=settings.API_PREFIX)
_UPLOAD_DIR = Path(settings.TEMP_DIR) / "api_uploads"
_UPLOAD_DIR.mkdir(parents=True, exist_ok=True)
_uploaded_files: dict[str, dict[str, str]] = {}
def _recognition_request_config_example() -> dict[str, Any]:
example: dict[str, Any] = {
"enable_speaker": True,
"match_speaker_registry": True,
"enable_text_cleanup": True,
"speaker_threshold": 0.6,
"hotwords": [
{"hotword": "阿里巴巴", "weight": 2.0},
{"hotword": "腾讯", "weight": 1.5},
],
}
if settings.ASR_ENABLE_WORD_TIMESTAMPS:
example["word_timestamps"] = False
return example
def _meeting_request_examples() -> dict[str, Any]:
return {
"url_with_hotwords": {
"summary": "URL + 结构化热词",
"description": "推荐的生产调用格式。",
"value": {
"model": "qwen3-asr-0.6b",
"audio_address": "https://example.com/media/meeting.mp4",
"config": _recognition_request_config_example(),
},
},
"local_path": {
"summary": "服务端本地路径",
"value": {
"audio_address": "/data/audio/meeting.m4a",
"config": {
"enable_speaker": True,
"match_speaker_registry": False,
"enable_text_cleanup": True,
},
},
},
}
def _recognition_request_config_schema_extra(schema: dict[str, Any], _model: Any) -> None:
schema["example"] = _recognition_request_config_example()
if settings.ASR_ENABLE_WORD_TIMESTAMPS:
return
properties = schema.get("properties")
if isinstance(properties, dict):
properties.pop("word_timestamps", None)
required = schema.get("required")
if isinstance(required, list):
schema["required"] = [item for item in required if item != "word_timestamps"]
class HotwordItem(BaseModel):
model_config = ConfigDict(json_schema_extra={"example": {"hotword": "汇智", "weight": 2.0}})
hotword: str = Field(
...,
title="热词内容",
description="需要优先匹配或纠正的业务词、专有名词、人名、产品名等。Qwen3-ASR 无原生热词能力,服务会在识别后做规则化热词纠正。",
examples=["通义千问"],
)
weight: float = Field(
1.0,
title="热词权重",
description="热词权重,兼容老项目 Model-Test-New 的传参格式;当前规则纠正主要使用热词文本,权重保留用于接口兼容和后续扩展。",
examples=[2.0],
)
class RecognitionRequestConfig(BaseModel):
model_config = ConfigDict(
json_schema_extra=_recognition_request_config_schema_extra
)
enable_speaker: bool = Field(
True,
title="是否启用说话人分离",
description="是否执行说话人分离。开启后返回的分段会带 speaker 字段;关闭后不做说话人分离和声纹库匹配。",
examples=[True],
)
match_speaker_registry: bool = Field(
True,
title="是否匹配已注册声纹库",
description="是否在说话人分离后匹配已注册声纹库。只有 enable_speaker=true 时生效。",
examples=[True],
)
enable_text_cleanup: bool = Field(
True,
title="是否启用文字去重",
description="是否启用识别结果去重/边界重叠裁剪,适合长音频分段识别后的重复文本清理。",
examples=[True],
)
speaker_threshold: Optional[float] = Field(
None,
title="声纹匹配阈值",
description="本次声纹库匹配阈值,范围通常为 0~1;不传则使用服务默认配置。",
examples=[0.6],
)
word_timestamps: bool = Field(
False,
title="是否返回字词级时间戳",
description="是否返回字词级时间戳。长会议场景建议默认关闭以降低耗时和结果体积。",
examples=[False],
)
hotwords: str | list[HotwordItem] = Field(
default_factory=list,
title="热词",
description=(
"热词参数。推荐传老项目兼容的结构化数组:"
"[{\"hotword\":\"通义千问\",\"weight\":2.0}];也兼容字符串:\"通义千问 2.0 ModelScope 1.5\"。"
"Qwen3-ASR 本身没有原生热词参数,服务会在识别后使用热词规则做文本纠正。"
),
examples=[[{"hotword": "通义千问", "weight": 2.0}, {"hotword": "ModelScope", "weight": 1.5}]],
)
sample_rate: int = Field(
16000,
title="采样率",
description="音频处理采样率,默认 16000。",
examples=[16000],
)
class TranscriptionCreateRequest(BaseModel):
model_config = ConfigDict(
json_schema_extra={
"example": {
"model": "qwen3-asr-0.6b",
"audio_address": "https://example.com/meeting.m4a",
"config": {
"enable_speaker": True,
"match_speaker_registry": True,
"enable_text_cleanup": True,
"speaker_threshold": 0.6,
"hotwords": [
{"hotword": "汇智", "weight": 2.0},
{"hotword": "通义千问", "weight": 2.5},
],
},
}
}
)
model: Optional[str] = Field(
None,
title="离线 ASR 模型 ID",
description="可选。不传则使用服务当前默认模型;可传当前可用的本地模型 ID,如 qwen3-asr-0.6b 或 qwen3-asr-1.7b。",
examples=["qwen3-asr-0.6b"],
max_length=128,
)
audio_address: str = Field(
...,
title="音视频地址",
description="必填。音频/视频文件地址,支持 HTTP/HTTPS URL、file:// 地址或服务端本地音视频路径。该接口生产调用不使用 file_id/file_url。",
examples=["https://example.com/meeting.m4a"],
)
config: RecognitionRequestConfig = Field(
default_factory=RecognitionRequestConfig,
title="识别配置",
description=(
"离线会议识别配置,包括说话人分离、声纹库匹配、文字去重和热词规则。"
if not settings.ASR_ENABLE_WORD_TIMESTAMPS
else "离线会议识别配置,包括说话人分离、声纹库匹配、文字去重、字词时间戳和热词规则。"
)
)
class SpeakerRegisterRequest(BaseModel):
model_config = ConfigDict(
json_schema_extra={"example": {"name": "张三", "user_id": "u001", "audio_address": "/data/speakers/zhangsan.wav"}}
)
name: str
audio_address: Optional[str] = None
file_url: Optional[str] = None
file_id: Optional[str] = None
user_id: Optional[str] = None
class SpeakerIdentifyRequest(BaseModel):
model_config = ConfigDict(json_schema_extra={"example": {"audio_address": "/data/audio/sample.wav", "threshold": 0.6}})
audio_address: Optional[str] = None
file_url: Optional[str] = None
file_id: Optional[str] = None
threshold: Optional[float] = None
def _get_meeting_transcription_description() -> str:
return """创建离线会议转写任务。该接口采用 **JSON 请求体**,适合业务系统直接提交会议音频/视频地址并异步轮询结果。
**音频输入方式:**
1. **HTTP/HTTPS URL**:`audio_address="https://example.com/meeting.mp4"`
2. **file:// 地址**:`audio_address="file:///data/audio/meeting.m4a"`
3. **服务端本地路径**:`audio_address="/data/audio/meeting.m4a"`
> 生产调用只使用 `audio_address`。`file_id` / `file_url` 不属于该接口参数,上传接口仅用于测试页面调试。
**请求参数:**
| 字段 | 类型 | 必填 | 默认值 | 说明 |
|------|------|------|--------|------|
| `model` | string/null | 否 | 服务默认模型 | 离线 ASR 模型 ID;不传则使用 `QWEN3_ASR_MODEL` 或自动选择结果 |
| `audio_address` | string | 是 | - | 音频/视频文件地址,支持 HTTP/HTTPS、`file://`、服务端本地路径 |
| `config` | object | 否 | 默认配置 | 识别配置对象 |
| `config.enable_speaker` | boolean | 否 | `true` | 是否启用说话人分离 |
| `config.match_speaker_registry` | boolean | 否 | `true` | 是否匹配已注册声纹库,仅在 `enable_speaker=true` 时生效 |
| `config.enable_text_cleanup` | boolean | 否 | `true` | 是否启用文字去重、边界重叠裁剪 |
| `config.speaker_threshold` | number/null | 否 | 服务默认值 | 本次声纹库匹配阈值,通常为 `0~1` |
""" + (
"| `config.word_timestamps` | boolean | 否 | `false` | 是否返回字词级时间戳;开启后会调用 Qwen3 ForcedAligner |\n"
if settings.ASR_ENABLE_WORD_TIMESTAMPS
else ""
) + """
| `config.hotwords` | array/string | 否 | `[]` | 热词。推荐老项目兼容数组格式:`[{\"hotword\":\"通义千问\",\"weight\":2.0}]` |
| `config.sample_rate` | integer | 否 | `16000` | 音频处理采样率 |
**热词说明:**
- Qwen3-ASR 本身没有原生热词参数。
- 本服务兼容 Model-Test-New 的热词格式,识别完成后通过规则做热词纠正。
- 推荐格式:`[{\"hotword\":\"通义千问\",\"weight\":2.0},{\"hotword\":\"ModelScope\",\"weight\":1.5}]`
- 兼容字符串格式:`\"通义千问 2.0 ModelScope 1.5\"`
**创建任务返回:**
| 字段 | 类型 | 说明 |
|------|------|------|
| `code` | integer | 业务状态码,`0` 表示成功 |
| `message` | string | 业务消息 |
| `data.task_id` | string | 任务 ID,用于查询识别进度与结果 |
| `data.status` | string | 初始状态,通常为 `queued` |
| `data.stage` | string | 当前阶段 |
| `data.percentage` | number | 当前进度百分比 |
**查询任务返回:**
调用 `GET /api/v1/asr/transcriptions/{task_id}` 查询任务。完成后 `data.result` 包含完整结果。
| 字段 | 类型 | 说明 |
|------|------|------|
| `data.status` | string | `queued` / `processing` / `completed` / `failed` |
| `data.result.text` | string | 完整转写文本 |
| `data.result.segments` | array | 分段结果 |
| `data.result.segments[].start` | number | 分段开始时间,秒 |
| `data.result.segments[].end` | number | 分段结束时间,秒 |
| `data.result.segments[].text` | string | 分段文本 |
| `data.result.segments[].speaker_id` | string | 说话人 ID,启用说话人分离时返回 |
| `data.result.segments[].speaker_name` | string | 匹配到声纹库时返回姓名,否则通常等于 speaker_id |
| `data.result.segments[].user_id` | string | 匹配到声纹库用户时返回 |
| `data.result.segments[].word_tokens` | array | 字词级时间戳,仅 `word_timestamps=true` 且模型支持时返回 |
| `data.result.audio_duration_seconds` | number | 音频总时长,秒 |
| `data.result.speaker_count` | integer | 说话人数量 |
| `data.result.speakers` | array | 说话人列表 |
"""
def _ok(data: Any) -> dict[str, Any]:
return {"code": 0, "message": "ok", "data": data or {}}
def _public_task_payload(task_id: str, task: dict[str, Any]) -> dict[str, Any]:
payload: dict[str, Any] = {"task_id": task_id, "status": task["status"]}
for key in (
"stage",
"message",
"percentage",
"detail",
"created_at",
"updated_at",
):
if task.get(key) is not None:
payload[key] = task[key]
if (
task["status"] == "processing"
and task.get("created_at") is not None
and task.get("percentage") not in (None, 0)
):
elapsed_seconds = max(1.0, time.time() - float(task["created_at"]))
percentage = max(1.0, float(task.get("percentage") or 0))
payload["eta_seconds"] = max(
0,
int(math.ceil(elapsed_seconds * (100.0 - percentage) / percentage)),
)
elif task["status"] == "completed":
payload["eta_seconds"] = 0
if task.get("result") is not None:
payload["result"] = task["result"]
if task.get("error"):
payload["error"] = task["error"]
return payload
def _require_db() -> None:
if not pg_speaker_db.is_connected:
raise HTTPException(
status_code=503,
detail={"code": 1001, "message": "speaker database is not connected"},
)
def _normalize_hotwords_for_asr(hotwords: str | list[HotwordItem] | None) -> str:
if isinstance(hotwords, str):
return hotwords.strip()
parts: list[str] = []
for item in hotwords or []:
text = item.hotword.strip()
if not text:
continue
parts.extend([text, f"{float(item.weight):g}"])
return " ".join(parts)
async def _download_url(url: str) -> bytes:
loop = asyncio.get_running_loop()
def _download() -> bytes:
response = requests.get(url, timeout=60)
response.raise_for_status()
return response.content
return await loop.run_in_executor(None, _download)
def _is_http_url(value: str) -> bool:
return value.lower().startswith(("http://", "https://"))
def _source_filename(source: str) -> str:
parsed = urlparse(source)
path = unquote(parsed.path or source)
return Path(path).name or "audio"
async def _read_local_file(path_value: str) -> bytes:
parsed = urlparse(path_value)
local_path = Path(unquote(parsed.path)) if parsed.scheme == "file" else Path(path_value).expanduser()
if not local_path.is_file():
raise HTTPException(status_code=404, detail={"code": 1001, "message": "local media file not found"})
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, local_path.read_bytes)
async def _resolve_source_bytes(audio_address: Optional[str], file_id: Optional[str]) -> tuple[bytes, Optional[str]]:
if audio_address:
source = audio_address.strip()
if _is_http_url(source):
return await _download_url(source), _source_filename(source)
return await _read_local_file(source), _source_filename(source)
if file_id:
item = _uploaded_files.get(file_id)
if not item:
matched_files = sorted(_UPLOAD_DIR.glob(f"{file_id}.*"))
if not matched_files:
raise HTTPException(status_code=404, detail={"code": 1001, "message": "file_id not found"})
matched_path = matched_files[0]
item = {"path": str(matched_path), "filename": matched_path.name}
_uploaded_files[file_id] = item
path = item["path"]
with open(path, "rb") as file_obj:
return file_obj.read(), item.get("filename")
raise HTTPException(status_code=400, detail={"code": 1001, "message": "audio_address or file_id is required"})
def _cleanup_uploaded_file(file_id: Optional[str]) -> None:
if not file_id:
return
item = _uploaded_files.pop(file_id, None)
candidate_paths: list[Path] = []
if item and item.get("path"):
candidate_paths.append(Path(item["path"]))
candidate_paths.extend(sorted(_UPLOAD_DIR.glob(f"{file_id}.*")))
for path in candidate_paths:
try:
path.unlink(missing_ok=True)
except Exception:
pass
def _format_segments(asr_result: Any) -> list[dict[str, Any]]:
segments: list[dict[str, Any]] = []
for segment in asr_result.segments:
item: dict[str, Any] = {
"start": round(float(segment.start_time), 2),
"end": round(float(segment.end_time), 2),
"duration": round(float(segment.end_time - segment.start_time), 2),
"text": segment.text,
}
if segment.speaker_id:
item["speaker_id"] = segment.speaker_id
if segment.speaker_name:
item["speaker_name"] = segment.speaker_name
elif segment.speaker_id:
item["speaker_name"] = segment.speaker_id
if segment.user_id:
item["user_id"] = segment.user_id
if segment.word_tokens:
item["word_tokens"] = [
{
"text": token.text,
"start": round(float(token.start_time), 3),
"end": round(float(token.end_time), 3),
}
for token in segment.word_tokens
]
segments.append(item)
return segments
def _build_transcription_result(asr_result: Any) -> dict[str, Any]:
segments = _format_segments(asr_result)
speakers: list[str] = []
for segment in segments:
name = str(segment.get("speaker_name") or segment.get("speaker_id") or "").strip()
if name and name not in speakers:
speakers.append(name)
return {
"text": asr_result.text,
"segments": segments,
"audio_duration_seconds": round(float(asr_result.duration), 2),
"speaker_count": len(speakers),
"speakers": speakers,
}
def _update_task_progress(
task_id: str,
*,
status: str,
stage: str,
message: str,
percentage: int,
detail: Optional[dict[str, Any]] = None,
) -> None:
current_task = get_task_record(task_id) or {}
current_status = str(current_task.get("status") or "")
current_percentage = int(current_task.get("percentage") or 0)
incoming_percentage = max(0, min(100, int(percentage)))
if (
current_status in {"queued", "processing"}
and status in {"queued", "processing"}
and incoming_percentage < current_percentage
):
return
payload: dict[str, Any] = {
"status": status,
"stage": stage,
"message": message,
"percentage": incoming_percentage,
}
if detail is not None:
payload["detail"] = detail
update_task_record(task_id, payload)
async def _run_transcription_task(
*,
task_id: str,
model: Optional[str],
audio_address: Optional[str],
config: RecognitionRequestConfig,
) -> None:
service = get_offline_transcription_service()
prepared_audio: Optional[PreparedAudio] = None
_update_task_progress(
task_id,
status="processing",
stage="preparing",
message="任务已开始,正在准备音频文件。",
percentage=5,
detail={"segment_total": 0, "segment_completed": 0},
)
try:
audio_data, filename = await _resolve_source_bytes(audio_address, None)
_update_task_progress(
task_id,
status="processing",
stage="normalizing",
message="音频来源已读取,正在转换为识别格式。",
percentage=8,
detail={"segment_total": 0, "segment_completed": 0},
)
prepared_audio = await service.prepare_upload(
audio_data=audio_data,
filename=filename,
task_id=task_id,
sample_rate=config.sample_rate or int(SampleRate.RATE_16000),
)
_update_task_progress(
task_id,
status="processing",
stage="transcribing",
message="音频已准备完成,正在执行离线识别。",
percentage=12,
detail={
"segment_total": 0,
"segment_completed": 0,
"audio_duration_seconds": round(float(prepared_audio.duration), 2),
},
)
def progress_callback(
stage: str,
message: str,
percentage: int,
detail: Optional[dict[str, Any]] = None,
) -> None:
_update_task_progress(
task_id,
status="processing",
stage=stage,
message=message,
percentage=percentage,
detail=detail,
)
asr_result = await service.transcribe(
prepared_audio,
OfflineTranscriptionOptions(
model_id=model,
sample_rate=config.sample_rate or int(SampleRate.RATE_16000),
hotwords=_normalize_hotwords_for_asr(config.hotwords),
enable_speaker_diarization=config.enable_speaker,
enable_speaker_identification=(
config.enable_speaker and config.match_speaker_registry
),
enable_text_cleanup=config.enable_text_cleanup,
word_timestamps=(
settings.ASR_ENABLE_WORD_TIMESTAMPS and config.word_timestamps
),
task_id=task_id,
progress_callback=progress_callback,
),
)
if (
config.enable_speaker
and config.match_speaker_registry
and config.speaker_threshold is not None
):
_update_task_progress(
task_id,
status="processing",
stage="speaker_matching",
message="正在使用本次阈值匹配已注册声纹库。",
percentage=95,
detail={
"segment_total": len(asr_result.segments),
"segment_completed": len(asr_result.segments),
"audio_duration_seconds": round(float(asr_result.duration), 2),
},
)
await get_speaker_registry_service().apply_registered_speakers(
asr_result,
threshold=config.speaker_threshold,
)
result = _build_transcription_result(asr_result)
_update_task_progress(
task_id,
status="completed",
stage="completed",
message="离线识别已经完成。",
percentage=100,
detail={
"segment_total": len(result.get("segments") or []),
"segment_completed": len(result.get("segments") or []),
"audio_duration_seconds": result.get("audio_duration_seconds"),
"speaker_count": result.get("speaker_count"),
"speakers": result.get("speakers"),
},
)
update_task_record(task_id, {"result": result})
except Exception as exc:
_update_task_progress(
task_id,
status="failed",
stage="failed",
message="离线识别执行失败。",
percentage=100,
)
update_task_record(task_id, {"error": str(exc)})
finally:
service.cleanup(prepared_audio)
@router.post("/files", tags=["会议离线接口"], summary="上传音频文件")
async def upload_file(file: UploadFile = File(...)) -> dict[str, Any]:
file_id = uuid.uuid4().hex
suffix = Path(file.filename or "").suffix or ".audio"
path = _UPLOAD_DIR / f"{file_id}{suffix}"
content = await file.read()
path.write_bytes(content)
_uploaded_files[file_id] = {"path": str(path), "filename": file.filename or path.name}
return _ok({"file_id": file_id, "file_path": str(path)})
@router.post(
"/asr/transcriptions",
tags=["会议离线接口"],
summary="创建离线会议识别任务",
description=_get_meeting_transcription_description(),
response_description="任务创建成功,返回 task_id;后续调用 GET /api/v1/asr/transcriptions/{task_id} 查询进度和结果。",
responses={
200: {
"description": "任务创建成功",
"content": {
"application/json": {
"example": {
"code": 0,
"message": "ok",
"data": {
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
"status": "queued",
"stage": "queued",
"message": "任务已创建,等待进入处理流程。",
"percentage": 0,
},
}
}
},
},
400: {
"description": "请求参数错误",
"content": {
"application/json": {
"example": {
"detail": {
"code": 1001,
"message": "audio_address is required",
}
}
}
},
},
},
)
async def create_transcription_task(
background_tasks: BackgroundTasks,
item: TranscriptionCreateRequest = Body(
...,
description="离线会议识别任务 JSON 请求体。必须提供 audio_address,可在 config 中配置说话人分离、声纹库匹配、文字去重、热词等。",
examples=_meeting_request_examples(),
),
) -> dict[str, Any]:
if not item.audio_address:
raise HTTPException(status_code=400, detail={"code": 1001, "message": "audio_address is required"})
try:
model_id = validate_offline_model_id(item.model)
except InvalidParameterException as exc:
raise HTTPException(status_code=400, detail={"code": 1001, "message": exc.message}) from exc
task_id = f"task_{uuid.uuid4().hex}"
create_task_record(
task_id,
{
"status": "queued",
"stage": "queued",
"message": "任务已创建,等待进入处理流程。",
"percentage": 0,
"result": None,
"detail": {"segment_total": 0, "segment_completed": 0},
},
)
background_tasks.add_task(
_run_transcription_task,
task_id=task_id,
model=model_id,
audio_address=item.audio_address,
config=item.config,
)
return _ok(
{
"task_id": task_id,
"status": "queued",
"stage": "queued",
"message": "任务已创建,等待进入处理流程。",
"percentage": 0,
}
)
@router.get(
"/asr/transcriptions/{task_id}",
tags=["会议离线接口"],
summary="查询离线会议识别任务",
description=(
"根据创建任务时返回的 `task_id` 查询离线会议识别进度和结果。"
"`status=completed` 时,`data.result` 会包含完整文本、分段、说话人和可选字词级时间戳。"
),
responses={
200: {
"description": "查询成功",
"content": {
"application/json": {
"examples": {
"processing": {
"summary": "处理中",
"value": {
"code": 0,
"message": "ok",
"data": {
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
"status": "processing",
"stage": "asr",
"message": "正在识别分段 80/127",
"percentage": 62,
"detail": {"segment_total": 127, "segment_completed": 80},
"eta_seconds": 38,
},
},
},
"completed": {
"summary": "已完成",
"value": {
"code": 0,
"message": "ok",
"data": {
"task_id": "task_6fd90daffd2f446f944b64fdf6c17572",
"status": "completed",
"stage": "completed",
"message": "识别完成。",
"percentage": 100,
"eta_seconds": 0,
"result": {
"text": "大家好,今天我们讨论通义千问和 ModelScope 的接入方案。",
"segments": [
{
"start": 0.0,
"end": 4.28,
"duration": 4.28,
"text": "大家好,今天我们讨论通义千问和 ModelScope 的接入方案。",
"speaker_id": "speaker_1",
"speaker_name": "张三",
"user_id": "u001",
"word_tokens": [
{"text": "大家好", "start": 0.12, "end": 0.68}
],
}
],
"audio_duration_seconds": 4669.05,
"speaker_count": 1,
"speakers": ["张三"],
},
},
},
},
}
}
},
},
404: {
"description": "任务不存在",
"content": {
"application/json": {
"example": {
"detail": {
"code": 1004,
"message": "task not found",
}
}
}
},
},
},
)
async def get_transcription_task(task_id: str) -> dict[str, Any]:
task = get_task_record(task_id, include_result=True)
if not task:
raise HTTPException(status_code=404, detail={"code": 1001, "message": "Task not found"})
return _ok(_public_task_payload(task_id, task))
@router.post("/speakers", tags=["声纹管理"], summary="注册声纹")
async def register_speaker(item: SpeakerRegisterRequest) -> dict[str, Any]:
_require_db()
temp_path: Optional[str] = None
try:
audio_data, filename = await _resolve_source_bytes(item.audio_address or item.file_url, item.file_id)
suffix = Path(filename or "").suffix or ".wav"
temp_path = await get_speaker_registry_service().save_upload_to_temp(audio_data, suffix=suffix)
result = await get_speaker_registry_service().register_file(
name=item.name,
file_path=temp_path,
user_id=item.user_id,
)
return _ok(result)
finally:
get_speaker_registry_service().cleanup_file(temp_path)
_cleanup_uploaded_file(item.file_id)
@router.get("/speakers", tags=["声纹管理"], summary="获取声纹人员列表")
async def list_speakers() -> dict[str, Any]:
_require_db()
return _ok(
{
"items": await get_speaker_registry_service().list_speakers(),
"speaker_model": "CampPlus",
}
)
@router.delete("/speakers/{speaker_id}", tags=["声纹管理"], summary="删除声纹人员")
async def delete_speaker(speaker_id: str) -> dict[str, Any]:
_require_db()
deleted = await get_speaker_registry_service().delete_speaker(speaker_id)
if not deleted:
raise HTTPException(status_code=404, detail={"code": 1001, "message": "Speaker not found"})
return _ok({"id": speaker_id})
@router.post("/speakers/identify", tags=["声纹管理"], summary="识别音频中的说话人")
async def identify_speaker(item: SpeakerIdentifyRequest) -> dict[str, Any]:
_require_db()
temp_path: Optional[str] = None
try:
audio_data, filename = await _resolve_source_bytes(item.audio_address or item.file_url, item.file_id)
suffix = Path(filename or "").suffix or ".wav"
temp_path = await get_speaker_registry_service().save_upload_to_temp(audio_data, suffix=suffix)
return _ok(
await get_speaker_registry_service().identify_file(
file_path=temp_path,
threshold=item.threshold,
)
)
finally:
get_speaker_registry_service().cleanup_file(temp_path)
_cleanup_uploaded_file(item.file_id)

View File

@ -0,0 +1,692 @@
# -*- coding: utf-8 -*-
"""
OpenAI 兼容 API
实现 OpenAI Audio API 规范,兼容 OpenAI SDK 和第三方客户端
"""
import asyncio
import json
import time
import logging
from typing import Any, Optional, List
from enum import Enum
from fastapi import APIRouter, File, Form, UploadFile, Request, HTTPException
from fastapi.responses import JSONResponse, PlainTextResponse, StreamingResponse
from pydantic import BaseModel, Field
from ...core.config import settings
from ...core.security import validate_openai_token
from ...core.exceptions import (
create_error_response,
InvalidParameterException,
)
from ...services.asr.model_selection import (
get_default_offline_model_id,
get_offline_model_ids,
validate_offline_model_id,
)
from ...services.asr.offline_transcription_service import (
OfflineTranscriptionOptions,
PreparedAudio,
get_offline_transcription_service,
)
logger = logging.getLogger(__name__)
def _parse_hidden_bool(raw_value: Any) -> bool:
if isinstance(raw_value, bool):
return raw_value
if raw_value is None:
return False
return str(raw_value).strip().lower() in {"1", "true", "yes", "on"}
router = APIRouter(prefix="/v1", tags=["OpenAI Compatible"])
HEARTBEAT_INTERVAL_SECONDS = 15.0
# ============= 枚举类型 =============
class ResponseFormat(str, Enum):
JSON = "json"
TEXT = "text"
SRT = "srt"
VERBOSE_JSON = "verbose_json"
VTT = "vtt"
# ============= 响应模型 =============
class TranscriptionSegment(BaseModel):
"""转写分段"""
id: int
seek: int = 0
start: float
end: float
text: str
tokens: List[int] = Field(default_factory=list)
temperature: float = 0.0
avg_logprob: float = 0.0
compression_ratio: float = 0.0
no_speech_prob: float = 0.0
speaker: Optional[str] = Field(default=None, description="说话人ID")
class TranscriptionWord(BaseModel):
"""转写词级别信息"""
word: str
start: float
end: float
class TranscriptionResponse(BaseModel):
"""简单转写响应 (json 格式)"""
text: str
class VerboseTranscriptionResponse(BaseModel):
"""详细转写响应 (verbose_json 格式)"""
task: str = "transcribe"
language: str
duration: float
text: str
segments: List[TranscriptionSegment] = Field(default_factory=list)
words: Optional[List[TranscriptionWord]] = None
class ModelObject(BaseModel):
"""模型对象"""
id: str
object: str = "model"
created: int = Field(default_factory=lambda: int(time.time()))
owned_by: str = "qwen3-asr"
class ModelsResponse(BaseModel):
"""模型列表响应"""
object: str = "list"
data: List[ModelObject]
# ============= 辅助函数 =============
def format_timestamp_srt(seconds: float) -> str:
"""格式化时间戳为 SRT 格式 (HH:MM:SS,mmm)"""
hours = int(seconds // 3600)
minutes = int((seconds % 3600) // 60)
secs = int(seconds % 60)
millis = int((seconds % 1) * 1000)
return f"{hours:02d}:{minutes:02d}:{secs:02d},{millis:03d}"
def format_timestamp_vtt(seconds: float) -> str:
"""格式化时间戳为 VTT 格式 (HH:MM:SS.mmm)"""
hours = int(seconds // 3600)
minutes = int((seconds % 3600) // 60)
secs = int(seconds % 60)
millis = int((seconds % 1) * 1000)
return f"{hours:02d}:{minutes:02d}:{secs:02d}.{millis:03d}"
def generate_srt(segments: List[TranscriptionSegment]) -> str:
"""生成 SRT 字幕格式"""
lines = []
for i, seg in enumerate(segments, 1):
start = format_timestamp_srt(seg.start)
end = format_timestamp_srt(seg.end)
lines.append(f"{i}")
lines.append(f"{start} --> {end}")
text = seg.text.strip()
if seg.speaker:
text = f"[{seg.speaker}] {text}"
lines.append(text)
lines.append("")
return "\n".join(lines)
def generate_vtt(segments: List[TranscriptionSegment]) -> str:
"""生成 WebVTT 字幕格式"""
lines = ["WEBVTT", ""]
for seg in segments:
start = format_timestamp_vtt(seg.start)
end = format_timestamp_vtt(seg.end)
lines.append(f"{start} --> {end}")
text = seg.text.strip()
if seg.speaker:
text = f"[{seg.speaker}] {text}"
lines.append(text)
lines.append("")
return "\n".join(lines)
def detect_language(text: str, language: Optional[str]) -> str:
"""检测识别语言。"""
if language:
return language
import re
if re.search(r"[\u4e00-\u9fff]", text):
return "zh"
return "en"
def build_transcription_payload(
*,
response_format: ResponseFormat,
asr_result,
audio_duration: float,
language: Optional[str],
) -> tuple[object, int, int]:
"""构建 OpenAI 转写响应载荷,并返回 segments / words 计数。"""
segments: List[TranscriptionSegment] = []
words: List[TranscriptionWord] = []
for i, seg in enumerate(asr_result.segments):
segments.append(
TranscriptionSegment(
id=i,
seek=int(seg.start_time * 100),
start=seg.start_time,
end=seg.end_time,
text=seg.text,
speaker=seg.speaker_id,
)
)
if seg.word_tokens:
for wt in seg.word_tokens:
words.append(
TranscriptionWord(
word=wt.text,
start=round(seg.start_time + wt.start_time, 3),
end=round(seg.start_time + wt.end_time, 3),
)
)
detected_language = detect_language(asr_result.text, language)
if response_format == ResponseFormat.VERBOSE_JSON:
payload = VerboseTranscriptionResponse(
task="transcribe",
language=detected_language,
duration=audio_duration,
text=asr_result.text,
segments=segments,
words=words if words else None,
).model_dump()
elif response_format == ResponseFormat.JSON:
payload = {"text": asr_result.text}
elif response_format == ResponseFormat.TEXT:
payload = asr_result.text
elif response_format == ResponseFormat.SRT:
if not segments:
segments = [
TranscriptionSegment(
id=0,
start=0,
end=audio_duration,
text=asr_result.text,
)
]
payload = generate_srt(segments)
elif response_format == ResponseFormat.VTT:
if not segments:
segments = [
TranscriptionSegment(
id=0,
start=0,
end=audio_duration,
text=asr_result.text,
)
]
payload = generate_vtt(segments)
else:
payload = {"text": asr_result.text}
return payload, len(segments), len(words)
def create_heartbeat_streaming_response(
*,
response_format: ResponseFormat,
inference_coro,
audio_duration: float,
language: Optional[str],
cleanup_callback,
) -> StreamingResponse:
"""为长耗时 JSON 响应生成带心跳的流式输出。"""
async def response_stream():
inference_task = asyncio.create_task(inference_coro)
heartbeat_count = 0
try:
while True:
done, _pending = await asyncio.wait(
{inference_task},
timeout=HEARTBEAT_INTERVAL_SECONDS,
return_when=asyncio.FIRST_COMPLETED,
)
if inference_task in done:
break
heartbeat_count += 1
logger.info(
"[OpenAI API] 发送响应心跳: "
f"format={response_format}, heartbeat_count={heartbeat_count}"
)
yield b" \n"
asr_result = await inference_task
logger.info(f"[OpenAI API] 识别完成: {len(asr_result.text)} 字符")
payload, segments_count, words_count = build_transcription_payload(
response_format=response_format,
asr_result=asr_result,
audio_duration=audio_duration,
language=language,
)
response_bytes = json.dumps(
payload,
ensure_ascii=False,
separators=(",", ":"),
).encode("utf-8")
logger.info(
"[OpenAI API] 准备发送 JSON 响应: "
f"format={response_format}, "
f"segments={segments_count}, "
f"words={words_count}, "
f"payload_bytes={len(response_bytes)}, "
f"heartbeat_count={heartbeat_count}"
)
yield response_bytes
finally:
cleanup_callback()
return StreamingResponse(
response_stream(),
media_type="application/json",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)
# ============= API 端点 =============
def _get_openai_model_description() -> str:
"""获取动态的模型描述"""
available_models = get_offline_model_ids()
default_model = get_default_offline_model_id()
model_descriptions = {
"qwen3-asr-1.7b": "Qwen3-ASR 1.7B,52 种语言,vLLM 高性能",
"qwen3-asr-0.6b": "Qwen3-ASR 0.6B,轻量版,适合小显存环境",
"qwen3-asr": "自动路由到当前已启动的 Qwen3-ASR 版本",
}
# 构建表格行
table_rows = []
for m in available_models:
desc = model_descriptions.get(m, "")
if m == default_model:
desc += "(默认)"
table_rows.append(f"| `{m}` | {desc} |")
return f"""返回当前可用的离线 Qwen3-ASR 模型列表(OpenAI `/v1/models` 兼容)。
**可用离线模型:**
| 模型 ID | 说明 |
|---------|------|
{chr(10).join(table_rows)}
**兼容性说明:**
- 支持 OpenAI SDK 和第三方客户端调用
- 当前默认模型根据显存自动选择;也可通过 `QWEN3_ASR_MODEL` 覆盖
"""
@router.get(
"/models",
response_model=ModelsResponse,
summary="列出可用离线模型",
description=_get_openai_model_description(),
)
async def list_models(request: Request):
"""列出可用离线模型 (OpenAI 兼容)"""
result, _ = validate_openai_token(request)
if not result:
response_data = create_error_response(
error_code="AUTHENTICATION_FAILED",
message="Invalid authentication",
)
return JSONResponse(content=response_data, status_code=401)
try:
# 使用动态模型列表
model_ids = get_offline_model_ids()
model_objects = []
for model_id in model_ids:
model_objects.append(ModelObject(
id=model_id,
owned_by="qwen3-asr",
))
return ModelsResponse(data=model_objects)
except Exception as e:
logger.error(f"获取模型列表失败: {e}")
raise HTTPException(status_code=500, detail=str(e))
def _get_transcription_description() -> str:
"""获取动态的转写端点描述"""
return f"""将音频文件转写为文本(完全兼容 OpenAI Audio API)。
**支持的音频格式与常见含音轨视频容器:**
`mp3`, `mp4`, `mpeg`, `mpga`, `m4a`, `wav`, `webm`, `flac`, `ogg`, `amr`, `pcm`, `mov`, `mkv`, `avi`
**音频输入方式:**
1. **文件上传**:通过 `file` 参数上传音频/视频文件(标准 OpenAI 方式)
2. **URL/本地路径读取**:通过 `audio_address` 参数提供音频/视频文件 URL(HTTP/HTTPS)或服务端本地路径
如果同时提供 `file` 和 `audio_address`,服务会优先使用 `file`,并忽略 `audio_address`。
**文件大小限制:**
- 最大支持 {settings.MAX_AUDIO_SIZE // (1024 * 1024)}MB(可通过 `MAX_AUDIO_SIZE` 环境变量配置)
- OpenAI 原生限制为 25MB
**说话人分离:**
- 默认开启 (`enable_speaker_diarization=true`)
- 启用后 `verbose_json` 格式的 segments 会包含 `speaker` 字段(如 "说话人1")
- 可设置 `enable_speaker_diarization=false` 关闭
**增强选项:**
- `enable_speaker_identification=true` 时会尝试匹配已注册声纹库,仅在说话人分离开启时生效
- `enable_text_cleanup=true` 时会执行文本去重、跨段重叠裁剪和口头语清理
- `hotwords` 可传热词字符串,格式:`热词1 权重1 热词2 权重2`
**输出格式:**
| 格式 | Content-Type | 说明 |
|------|-------------|------|
| `json` | application/json | 简单 JSON,仅含 text 字段(默认) |
| `text` | text/plain | 纯文本 |
| `verbose_json` | application/json | 详细 JSON,含时间戳、分段和说话人 |
| `srt` | text/plain | SRT 字幕格式 |
| `vtt` | text/vtt | WebVTT 字幕格式 |
**模型选择:**
- 默认使用当前服务启用的 Qwen3-ASR 模型
- 可通过可选 `model` 表单字段指定当前可用离线模型
- 通过 `QWEN3_ASR_MODEL` 控制服务端默认模型型号
- `/v1/models` 仍可用于查看当前服务端实际在线模型
**兼容参数:**
`prompt`、`temperature`、`timestamp_granularities` 参数已保留;其中 `prompt` 可作为热词提示的兼容入口
"""
@router.post(
"/audio/transcriptions",
summary="音频转写",
description=_get_transcription_description(),
responses={
200: {
"description": "转写成功",
"content": {
"application/json": {
"example": {"text": "今天天气不错,明天可能会下雨。"}
},
"text/plain": {
"example": "今天天气不错,明天可能会下雨。"
},
},
},
400: {
"description": "请求错误",
"content": {
"application/json": {
"example": {
"error_code": "INVALID_PARAMETER",
"message": f"File too large. Maximum size is {settings.MAX_AUDIO_SIZE // (1024 * 1024)}MB",
"task_id": "",
"timestamp": "2025-01-31T12:00:00Z",
"details": {}
}
}
},
},
401: {
"description": "认证失败",
"content": {
"application/json": {
"example": {
"error_code": "AUTHENTICATION_FAILED",
"message": "Invalid API key",
"task_id": "",
"timestamp": "2025-01-31T12:00:00Z",
"details": {}
}
}
},
},
},
)
async def create_transcription(
request: Request,
# 1. 音频输入(二选一)
file: Optional[UploadFile] = File(
default=None,
description="要转写的音频/视频文件。若同时提供 audio_address,服务会优先使用这里上传的文件"
),
audio_address: Optional[str] = Form(
default=None,
description="音频/视频文件地址,支持 HTTP/HTTPS URL、file:// 或服务端本地路径。仅当 file 为空时使用;若同时上传 file,服务会忽略此参数",
json_schema_extra={"example": "https://media.cdn.vect.one/podcast_demo.mp4"},
),
model: Optional[str] = Form(
default=None,
description="可选。离线 ASR 模型 ID;不传则使用服务当前默认模型",
examples=["qwen3-asr-0.6b"],
),
language: Optional[str] = Form(
None,
description="音频语言代码(ISO-639-1),如 zh/en/ja,不填则自动检测",
examples=["zh", "en", "ja"],
),
# 4. 功能开关
enable_speaker_diarization: bool = Form(
True,
description="是否启用说话人分离(默认开启)。启用后响应 segments 会包含 speaker 字段"
),
enable_speaker_identification: bool = Form(
True,
description="是否匹配已注册声纹库。仅在 enable_speaker_diarization=true 时生效"
),
enable_text_cleanup: bool = Form(
True,
description="是否启用文本去重、跨段重叠裁剪和口头语清理"
),
hotwords: Optional[str] = Form(
None,
description="热词字符串,格式:热词1 权重1 热词2 权重2"
),
# 5. 输出选项
response_format: ResponseFormat = Form(
ResponseFormat.VERBOSE_JSON,
description="输出格式",
examples=["verbose_json", "json", "text", "srt", "vtt"],
),
# 6. 兼容性参数(暂不支持)
prompt: Optional[str] = Form(None, description="提示文本(暂不支持,保留兼容)"), # noqa: ARG001
temperature: Optional[float] = Form(0, description="采样温度(暂不支持,保留兼容)"), # noqa: ARG001
timestamp_granularities: Optional[List[str]] = Form( # noqa: ARG001
None,
alias="timestamp_granularities[]",
description="时间戳粒度(暂不支持,保留兼容)"
),
):
"""音频转写 API (OpenAI Audio API 兼容)"""
# 标记暂不支持的参数(保留以兼容 OpenAI API)
_ = (temperature, timestamp_granularities)
hotword_text = (hotwords or prompt or "").strip()
form_data = await request.form()
requested_word_timestamps = _parse_hidden_bool(form_data.get("word_timestamps"))
word_timestamps = settings.ASR_ENABLE_WORD_TIMESTAMPS and requested_word_timestamps
prepared_audio: Optional[PreparedAudio] = None
response_cleanup_managed = False
logger.info(f"[OpenAI API] 收到转写请求: model={model or 'default'}, format={response_format}, "
f"speaker_diarization={enable_speaker_diarization}, word_level={word_timestamps}, "
f"audio_address={'有' if audio_address else '无'}")
# 验证输入:至少提供一种输入源;若二者同时存在,优先 file
if not file and not audio_address:
response_data = create_error_response(
error_code="INVALID_PARAMETER",
message="必须提供 file(上传文件)或 audio_address(音频 URL)其中之一",
)
return JSONResponse(content=response_data, status_code=400)
transcription_service = get_offline_transcription_service()
try:
result, _ = validate_openai_token(request)
if not result:
response_data = create_error_response(
error_code="AUTHENTICATION_FAILED",
message="Invalid authentication",
)
return JSONResponse(content=response_data, status_code=401)
model_id = validate_offline_model_id(model)
# 处理音频输入:优先 file,其次 audio_address
if file is not None:
if audio_address:
logger.info("[OpenAI API] 检测到同时提供 file 和 audio_address,已忽略 audio_address")
logger.info(f"[OpenAI API] 从上传文件读取音频: {file.filename}")
audio_data = await file.read()
prepared_audio = await transcription_service.prepare_upload(
audio_data=audio_data,
filename=file.filename if file else None,
task_id=f"openai-{int(time.time() * 1000)}",
sample_rate=16000,
)
else:
logger.info(f"[OpenAI API] 从 URL 下载音频: {audio_address}")
prepared_audio = await transcription_service.prepare_from_request(
request=request,
audio_address=audio_address,
task_id=f"openai-{int(time.time() * 1000)}",
sample_rate=16000,
)
inference_coro = transcription_service.transcribe(
prepared_audio,
OfflineTranscriptionOptions(
model_id=model_id,
sample_rate=16000,
hotwords=hotword_text,
enable_speaker_diarization=enable_speaker_diarization,
enable_speaker_identification=(
enable_speaker_diarization and enable_speaker_identification
),
enable_text_cleanup=enable_text_cleanup,
word_timestamps=word_timestamps,
),
)
audio_duration = prepared_audio.duration
# 根据 response_format 返回不同格式
if response_format == ResponseFormat.TEXT:
asr_result = await inference_coro
logger.info(f"[OpenAI API] 识别完成: {len(asr_result.text)} 字符")
payload, _, _ = build_transcription_payload(
response_format=response_format,
asr_result=asr_result,
audio_duration=audio_duration,
language=language,
)
return PlainTextResponse(content=payload)
elif response_format == ResponseFormat.SRT:
asr_result = await inference_coro
logger.info(f"[OpenAI API] 识别完成: {len(asr_result.text)} 字符")
payload, _, _ = build_transcription_payload(
response_format=response_format,
asr_result=asr_result,
audio_duration=audio_duration,
language=language,
)
return PlainTextResponse(content=payload, media_type="text/plain")
elif response_format == ResponseFormat.VTT:
asr_result = await inference_coro
logger.info(f"[OpenAI API] 识别完成: {len(asr_result.text)} 字符")
payload, _, _ = build_transcription_payload(
response_format=response_format,
asr_result=asr_result,
audio_duration=audio_duration,
language=language,
)
return PlainTextResponse(content=payload, media_type="text/vtt")
elif response_format in {ResponseFormat.VERBOSE_JSON, ResponseFormat.JSON}:
response_cleanup_managed = True
return create_heartbeat_streaming_response(
response_format=response_format,
inference_coro=inference_coro,
audio_duration=audio_duration,
language=language,
cleanup_callback=lambda: transcription_service.cleanup(prepared_audio),
)
else:
asr_result = await inference_coro
logger.info(f"[OpenAI API] 识别完成: {len(asr_result.text)} 字符")
payload, _, _ = build_transcription_payload(
response_format=ResponseFormat.JSON,
asr_result=asr_result,
audio_duration=audio_duration,
language=language,
)
return JSONResponse(content=payload)
except HTTPException as http_exc:
# 将 HTTPException 转换为标准错误格式
logger.error(f"[OpenAI API] HTTP异常: {http_exc.detail}")
response_data = create_error_response(
error_code="DEFAULT_CLIENT_ERROR" if http_exc.status_code < 500 else "DEFAULT_SERVER_ERROR",
message=http_exc.detail,
)
return JSONResponse(content=response_data, status_code=http_exc.status_code)
except InvalidParameterException as e:
logger.error(f"[OpenAI API] 参数异常: {e.message}")
response_data = create_error_response(
error_code=e.error_code,
message=e.message,
details=e.details,
)
return JSONResponse(content=response_data, status_code=400)
except Exception as e:
logger.error(f"[OpenAI API] 转写失败: {e}")
# 使用标准错误格式
response_data = create_error_response(
error_code="DEFAULT_SERVER_ERROR",
message=str(e),
)
return JSONResponse(content=response_data, status_code=500)
finally:
if not response_cleanup_managed:
transcription_service.cleanup(prepared_audio)

View File

@ -0,0 +1,47 @@
# -*- coding: utf-8 -*-
"""WebSocket ASR API routes."""
import logging
import time
import uuid
from typing import Optional
from fastapi import APIRouter, WebSocket
from ...services.qwen3_websocket_asr import Qwen3ASRService
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/ws/v1/asr", tags=["WebSocket ASR"])
@router.websocket("/funasr")
async def funasr_websocket(websocket: WebSocket) -> None:
await websocket.accept()
task_id = f"deprecated_funasr_ws_{int(time.time())}_{id(websocket)}"
try:
await websocket.send_json(
{
"type": "error",
"task_id": task_id,
"code": "FUNASR_REALTIME_REMOVED",
"message": "FunASR/Paraformer realtime websocket has been removed. Use /ws/v1/asr/qwen instead.",
}
)
except Exception:
pass
finally:
await websocket.close(code=1008, reason="Use /ws/v1/asr/qwen")
_qwen3_service = Qwen3ASRService()
@router.websocket("")
@router.websocket("/qwen")
async def qwen_asr_websocket(
websocket: WebSocket,
task_id: Optional[str] = None,
) -> None:
if task_id is None:
task_id = str(uuid.uuid4())[:8]
await _qwen3_service.handle_connection(websocket, task_id)

58
app/bootstrap.py 100644
View File

@ -0,0 +1,58 @@
# -*- coding: utf-8 -*-
"""Shared bootstrap helpers for process startup."""
from __future__ import annotations
import sys
def ensure_models_downloaded(interactive: bool) -> bool:
"""Ensure declared deployment models exist locally, downloading if needed."""
try:
from app.utils.download_models import check_all_models, download_models
missing = check_all_models()
if not missing:
return True
print(f"\n⚠️ 检测到 {len(missing)} 个模型未下载")
for model_id, *_ in missing:
print(f" - {model_id}")
print("\n将自动下载缺失模型后继续启动。")
if download_models(auto_mode=True):
return True
print("\n模型自动下载失败。")
if interactive:
print("可手动运行以下命令排查:")
print(" uv run python -m app.utils.download_models")
print(" ./scripts/prepare-models.sh")
else:
print("非交互式终端下请确认网络可用,或预先准备模型缓存。")
return False
except Exception as exc:
print(f"⚠️ 模型检查失败: {exc}")
return False
def run_cli_preflight() -> bool:
"""Preflight checks for the CLI entrypoint."""
try:
from app.core.accelerator import get_accelerator_info, validate_accelerator_runtime
ok, message = validate_accelerator_runtime()
info = get_accelerator_info()
print(
"Accelerator | "
f"vendor={info.vendor} runtime={info.runtime} device={info.device} "
f"count={info.device_count}"
)
if not ok:
print(f"加速器运行栈检查失败: {message}")
return False
except Exception as exc:
print(f"加速器检查失败: {exc}")
return False
return ensure_models_downloaded(interactive=sys.stdin.isatty())

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
核心模块
包含配置、异常、安全等基础组件
"""

View File

@ -0,0 +1,162 @@
# -*- coding: utf-8 -*-
"""Unified accelerator detection and runtime validation."""
from __future__ import annotations
import os
from functools import lru_cache
from typing import Optional
from .accelerators import (
AcceleratorInfo,
IluvatarAcceleratorAdapter,
MetaxAcceleratorAdapter,
MThreadsAcceleratorAdapter,
NvidiaAcceleratorAdapter,
)
SUPPORTED_ACCELERATORS = {"auto", "cpu", "nvidia", "metax", "iluvatar", "mthreads"}
def _cpu_info(reason: str = "") -> AcceleratorInfo:
return AcceleratorInfo(
vendor="cpu",
runtime="cpu",
device="cpu",
available=True,
reason=reason,
)
def _normalize_configured(configured: Optional[str]) -> str:
value = (configured or os.getenv("ACCELERATOR") or "auto").strip().lower()
if value == "cuda":
return "nvidia"
if value in {"maca", "muxi", "mx"}:
return "metax"
if value in {"ix", "tianshu", "天数"}:
return "iluvatar"
if value in {"mthreads", "musa", "moorethreads", "摩尔线程"}:
return "mthreads"
if value not in SUPPORTED_ACCELERATORS:
raise ValueError(
f"Unsupported ACCELERATOR={value!r}; expected one of: "
f"{', '.join(sorted(SUPPORTED_ACCELERATORS))}"
)
return value
def _detect_uncached(configured: Optional[str]) -> AcceleratorInfo:
accelerator = _normalize_configured(configured)
if accelerator == "cpu":
return _cpu_info("forced by ACCELERATOR=cpu")
if accelerator == "nvidia":
return NvidiaAcceleratorAdapter().detect()
if accelerator == "metax":
return MetaxAcceleratorAdapter().detect()
if accelerator == "iluvatar":
return IluvatarAcceleratorAdapter().detect()
if accelerator == "mthreads":
return MThreadsAcceleratorAdapter().detect()
# Auto mode prefers vendor SMI commands before NVIDIA. This avoids a vendor
# PyTorch build exposing torch.cuda and being mistaken for NVIDIA.
metax = MetaxAcceleratorAdapter().detect()
if metax.available:
return metax
iluvatar = IluvatarAcceleratorAdapter().detect()
if iluvatar.available:
return iluvatar
mthreads = MThreadsAcceleratorAdapter().detect()
if mthreads.available:
return mthreads
nvidia = NvidiaAcceleratorAdapter().detect()
if nvidia.available:
return nvidia
return _cpu_info("no supported accelerator detected")
@lru_cache(maxsize=8)
def _detect_cached(configured: str) -> AcceleratorInfo:
return _detect_uncached(configured)
def detect_accelerator(configured: Optional[str] = None, *, refresh: bool = False) -> AcceleratorInfo:
"""Detect the active accelerator.
Args:
configured: Optional override matching ACCELERATOR values.
refresh: Clear cached detection before probing.
"""
normalized = _normalize_configured(configured)
if refresh:
_detect_cached.cache_clear()
return _detect_cached(normalized)
def get_accelerator_info(*, refresh: bool = False) -> AcceleratorInfo:
try:
from app.core.config import settings
configured = settings.ACCELERATOR
except Exception:
configured = os.getenv("ACCELERATOR", "auto")
return detect_accelerator(configured, refresh=refresh)
def validate_accelerator_runtime() -> tuple[bool, str]:
"""Validate explicit accelerator selections before model loading."""
try:
from app.core.config import settings
configured = _normalize_configured(settings.ACCELERATOR)
except Exception:
configured = _normalize_configured(os.getenv("ACCELERATOR", "auto"))
info = detect_accelerator(configured, refresh=True)
if configured in {"auto", "cpu", ""}:
return True, ""
if not info.available:
return False, f"ACCELERATOR={configured} requested but {info.reason}"
if configured == "metax" and not info.metadata.get("torch_cuda_available"):
return (
False,
"ACCELERATOR=metax detected mx-smi devices, but the active Python "
"environment does not expose torch.cuda. Run ./scripts/sync_metax_env.sh "
"or install the MetaX MACA PyTorch stack.",
)
if configured == "iluvatar" and not info.metadata.get("torch_cuda_available"):
return (
False,
"ACCELERATOR=iluvatar detected ixsmi devices, but the active Python "
"environment does not expose torch.cuda. Use the official Iluvatar "
"vLLM image as the base image or install the Iluvatar PyTorch stack.",
)
if configured == "mthreads" and not info.metadata.get("torch_cuda_available"):
return (
False,
"ACCELERATOR=mthreads detected mthreads-gmi devices, but the active Python "
"environment does not expose torch.cuda. Use the official Moore Threads "
"MUSA vLLM image as the base image or install the matching MUSA PyTorch stack.",
)
if configured == "nvidia" and not info.metadata.get("torch_cuda_available"):
return (
False,
"ACCELERATOR=nvidia detected NVIDIA devices, but torch.cuda is not "
"available in the active Python environment. Run ./scripts/sync_gpu_env.sh.",
)
return True, ""

View File

@ -0,0 +1,16 @@
# -*- coding: utf-8 -*-
"""Accelerator adapter exports."""
from .base import AcceleratorInfo
from .iluvatar import IluvatarAcceleratorAdapter
from .metax import MetaxAcceleratorAdapter
from .mthreads import MThreadsAcceleratorAdapter
from .nvidia import NvidiaAcceleratorAdapter
__all__ = [
"AcceleratorInfo",
"IluvatarAcceleratorAdapter",
"MetaxAcceleratorAdapter",
"MThreadsAcceleratorAdapter",
"NvidiaAcceleratorAdapter",
]

View File

@ -0,0 +1,109 @@
# -*- coding: utf-8 -*-
"""Shared accelerator adapter primitives."""
from __future__ import annotations
import os
import re
import shutil
import subprocess
from dataclasses import dataclass, field
from typing import Optional, Protocol
@dataclass(frozen=True)
class AcceleratorInfo:
"""Normalized hardware/runtime information used by the application."""
vendor: str
runtime: str
device: str
device_count: int = 0
visible_devices: tuple[str, ...] = ()
total_memory_gb: float = 0.0
smi_command: Optional[str] = None
available: bool = False
reason: str = ""
metadata: dict[str, object] = field(default_factory=dict)
@property
def is_gpu(self) -> bool:
return self.vendor not in {"cpu", "unknown"} and self.available
def as_dict(self) -> dict[str, object]:
return {
"vendor": self.vendor,
"runtime": self.runtime,
"device": self.device,
"device_count": self.device_count,
"visible_devices": list(self.visible_devices),
"total_memory_gb": self.total_memory_gb,
"smi_command": self.smi_command,
"available": self.available,
"reason": self.reason,
"metadata": self.metadata,
}
@property
def supports_sharded(self) -> bool:
value = self.metadata.get("supports_sharded")
return bool(value)
class AcceleratorAdapter(Protocol):
vendor: str
runtime: str
def detect(self) -> AcceleratorInfo:
"""Return normalized accelerator info."""
def command_exists(command: str) -> bool:
return shutil.which(command) is not None
def run_command(command: list[str], timeout: float = 3.0) -> str:
try:
completed = subprocess.run(
command,
check=False,
capture_output=True,
text=True,
timeout=timeout,
)
except (OSError, subprocess.TimeoutExpired):
return ""
if completed.returncode != 0:
return ""
return completed.stdout.strip()
def parse_visible_devices(raw: str | None) -> tuple[str, ...]:
value = (raw or "").strip()
if not value or value.lower() in {"all", "none", "void"}:
return ()
return tuple(part.strip() for part in value.split(",") if part.strip())
def first_env_devices(names: tuple[str, ...]) -> tuple[str, ...]:
shared = parse_visible_devices(os.getenv("ASR_VISIBLE_DEVICES"))
if shared:
return shared
for name in names:
devices = parse_visible_devices(os.getenv(name))
if devices:
return devices
return ()
def parse_memory_gb(text: str) -> float:
"""Parse the smallest memory value from common smi outputs."""
values: list[float] = []
for number, unit in re.findall(r"([0-9]+(?:\.[0-9]+)?)\s*(GiB|GB|MiB|MB)", text, re.I):
value = float(number)
normalized_unit = unit.lower()
if normalized_unit in {"mib", "mb"}:
value = value / 1024
values.append(value)
return min(values) if values else 0.0

View File

@ -0,0 +1,106 @@
# -*- coding: utf-8 -*-
"""Iluvatar/Tianshu accelerator adapter."""
from __future__ import annotations
import re
from .base import (
AcceleratorInfo,
command_exists,
first_env_devices,
parse_memory_gb,
run_command,
)
class IluvatarAcceleratorAdapter:
vendor = "iluvatar"
runtime = "ix"
smi_command = "ixsmi"
visible_env_names = (
"ILUVATAR_VISIBLE_DEVICES",
"IX_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
)
def _query_device_count(self) -> tuple[int, str]:
list_output = run_command([self.smi_command, "-L"])
if list_output:
lines = [line for line in list_output.splitlines() if line.strip()]
gpu_lines = [
line
for line in lines
if re.search(r"\b(gpu|device|card)\b", line, re.I)
]
return len(gpu_lines or lines), list_output
table_output = run_command([self.smi_command])
if table_output:
indexes = set(re.findall(r"(?:GPU|Device|Card)\s*[:#]?\s*([0-9]+)", table_output, re.I))
if not indexes:
indexes = set(re.findall(r"^\s*\|\s*([0-9]+)\s+", table_output, re.M))
if indexes:
return len(indexes), table_output
if re.search(r"\bIluvatar\b|\b天数\b|\bIX\b", table_output, re.I):
return 1, table_output
return 0, ""
def detect(self) -> AcceleratorInfo:
visible_devices = first_env_devices(self.visible_env_names)
smi_available = command_exists(self.smi_command)
if not smi_available:
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cpu",
visible_devices=visible_devices,
smi_command=None,
available=False,
reason="ixsmi not found",
)
device_count, raw_output = self._query_device_count()
if visible_devices:
device_count = min(device_count or len(visible_devices), len(visible_devices))
torch_cuda_available = False
torch_count = 0
torch_memory_gb = 0.0
torch_version = ""
try:
import torch
torch_version = getattr(torch, "__version__", "")
torch_cuda_available = bool(torch.cuda.is_available())
if torch_cuda_available:
torch_count = int(torch.cuda.device_count())
torch_memory_gb = min(
torch.cuda.get_device_properties(i).total_memory / (1024**3)
for i in range(torch_count)
)
except Exception:
pass
if torch_count:
device_count = min(torch_count, len(visible_devices)) if visible_devices else torch_count
total_memory_gb = torch_memory_gb or parse_memory_gb(raw_output)
available = device_count > 0
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cuda:0" if available else "cpu",
device_count=device_count,
visible_devices=visible_devices,
total_memory_gb=total_memory_gb,
smi_command=self.smi_command,
available=available,
reason="" if available else "ixsmi did not report any device",
metadata={
"torch_cuda_available": torch_cuda_available,
"torch_version": torch_version,
"supports_sharded": available,
},
)

View File

@ -0,0 +1,174 @@
# -*- coding: utf-8 -*-
"""MetaX/MuXi MACA accelerator adapter."""
from __future__ import annotations
import json
import re
from typing import Any
from .base import (
AcceleratorInfo,
command_exists,
first_env_devices,
parse_memory_gb,
run_command,
)
class MetaxAcceleratorAdapter:
vendor = "metax"
runtime = "maca"
smi_command = "mx-smi"
visible_env_names = (
"METAX_VISIBLE_DEVICES",
"MACA_VISIBLE_DEVICES",
"MX_VISIBLE_DEVICES",
)
def _json_output(self, command: list[str]) -> Any:
output = run_command(command)
if not output:
return None
try:
return json.loads(output)
except json.JSONDecodeError:
return None
def _walk_json(self, value: Any):
if isinstance(value, dict):
yield value
for child in value.values():
yield from self._walk_json(child)
elif isinstance(value, list):
for item in value:
yield from self._walk_json(item)
def _count_from_json(self, value: Any) -> int:
max_index = -1
device_like = 0
for item in self._walk_json(value):
keys = {str(key).lower(): key for key in item.keys()}
if any(key in keys for key in ("gpu id", "gpu_id", "gpu", "device id", "device_id", "index")):
device_like += 1
for key_name in ("gpu id", "gpu_id", "device id", "device_id", "index", "id"):
original = keys.get(key_name)
if original is None:
continue
try:
max_index = max(max_index, int(item[original]))
except (TypeError, ValueError):
continue
if max_index >= 0:
return max_index + 1
return device_like
def _memory_from_json(self, value: Any) -> float:
memory_values: list[float] = []
for item in self._walk_json(value):
for raw_key, raw_value in item.items():
key = str(raw_key).lower()
if not any(token in key for token in ("memory", "mem", "hbm", "vram", "容量")):
continue
if isinstance(raw_value, (int, float)):
numeric = float(raw_value)
# mx-smi reports are commonly MiB for memory counters.
if numeric > 1024:
numeric = numeric / 1024
memory_values.append(numeric)
continue
if isinstance(raw_value, str):
parsed = parse_memory_gb(raw_value)
if parsed > 0:
memory_values.append(parsed)
return min(memory_values) if memory_values else 0.0
def _query_device_count(self) -> tuple[int, str]:
json_info = self._json_output([self.smi_command, "-j"])
json_count = self._count_from_json(json_info)
if json_count > 0:
return json_count, json.dumps(json_info, ensure_ascii=False)
list_output = run_command([self.smi_command, "-L"])
if list_output:
lines = [line for line in list_output.splitlines() if line.strip()]
gpu_lines = [
line
for line in lines
if re.search(r"\b(gpu|device|card)\b", line, re.I)
]
return len(gpu_lines or lines), list_output
table_output = run_command([self.smi_command])
if table_output:
indexes = set(re.findall(r"(?:GPU|Device|Card)\s*[:#]?\s*([0-9]+)", table_output, re.I))
if not indexes:
indexes = set(re.findall(r"^\s*\|\s*([0-9]+)\s+", table_output, re.M))
return len(indexes), table_output
return 0, ""
def _query_memory_gb(self, fallback_output: str) -> float:
for command in (
[self.smi_command, "--show-memory", "-j"],
[self.smi_command, "--show-hwinfo", "-j"],
[self.smi_command, "-j"],
):
parsed = self._memory_from_json(self._json_output(command))
if parsed > 0:
return parsed
memory_output = run_command([self.smi_command, "--show-memory"])
hwinfo_output = run_command([self.smi_command, "--show-hwinfo"])
return (
parse_memory_gb(memory_output)
or parse_memory_gb(hwinfo_output)
or parse_memory_gb(fallback_output)
)
def detect(self) -> AcceleratorInfo:
visible_devices = first_env_devices(self.visible_env_names)
smi_available = command_exists(self.smi_command)
if not smi_available:
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cpu",
visible_devices=visible_devices,
smi_command=None,
available=False,
reason="mx-smi not found",
)
device_count, raw_output = self._query_device_count()
if visible_devices:
device_count = min(device_count or len(visible_devices), len(visible_devices))
total_memory_gb = self._query_memory_gb(raw_output)
torch_cuda_available = False
torch_version = ""
try:
import torch
torch_version = getattr(torch, "__version__", "")
torch_cuda_available = bool(torch.cuda.is_available())
except Exception:
pass
available = device_count > 0
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cuda:0" if available else "cpu",
device_count=device_count,
visible_devices=visible_devices,
total_memory_gb=total_memory_gb,
smi_command=self.smi_command,
available=available,
reason="" if available else "mx-smi did not report any device",
metadata={
"torch_cuda_available": torch_cuda_available,
"torch_version": torch_version,
"supports_sharded": available,
},
)

View File

@ -0,0 +1,114 @@
# -*- coding: utf-8 -*-
"""Moore Threads / MUSA accelerator adapter."""
from __future__ import annotations
import re
from .base import (
AcceleratorInfo,
command_exists,
first_env_devices,
parse_memory_gb,
run_command,
)
class MThreadsAcceleratorAdapter:
vendor = "mthreads"
runtime = "musa"
smi_command = "mthreads-gmi"
visible_env_names = (
"MTHREADS_VISIBLE_DEVICES",
"MUSA_VISIBLE_DEVICES",
"CUDA_VISIBLE_DEVICES",
)
def _query_device_count(self) -> tuple[int, str]:
for command in (
[self.smi_command, "-L"],
[self.smi_command, "list"],
[self.smi_command],
):
output = run_command(command)
if not output:
continue
lines = [line for line in output.splitlines() if line.strip()]
gpu_lines = [
line
for line in lines
if re.search(r"\b(gpu|device|card|musa|mthreads|moore)\b", line, re.I)
]
if gpu_lines:
return len(gpu_lines), output
indexes = set(
re.findall(r"(?:GPU|Device|Card)\s*[:#]?\s*([0-9]+)", output, re.I)
)
if not indexes:
indexes = set(re.findall(r"^\s*\|\s*([0-9]+)\s+", output, re.M))
if indexes:
return len(indexes), output
if re.search(r"\bMUSA\b|\bMoore\s+Threads\b|\bMThreads\b", output, re.I):
return max(len(lines), 1), output
return 0, ""
def detect(self) -> AcceleratorInfo:
visible_devices = first_env_devices(self.visible_env_names)
smi_available = command_exists(self.smi_command)
if not smi_available:
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cpu",
visible_devices=visible_devices,
smi_command=None,
available=False,
reason="mthreads-gmi not found",
)
device_count, raw_output = self._query_device_count()
if visible_devices:
device_count = min(device_count or len(visible_devices), len(visible_devices))
torch_cuda_available = False
torch_count = 0
torch_memory_gb = 0.0
torch_version = ""
try:
import torch
torch_version = getattr(torch, "__version__", "")
torch_cuda_available = bool(torch.cuda.is_available())
if torch_cuda_available:
torch_count = int(torch.cuda.device_count())
torch_memory_gb = min(
torch.cuda.get_device_properties(i).total_memory / (1024**3)
for i in range(torch_count)
)
except Exception:
pass
if torch_count:
device_count = min(torch_count, len(visible_devices)) if visible_devices else torch_count
total_memory_gb = torch_memory_gb or parse_memory_gb(raw_output)
available = device_count > 0
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cuda:0" if available else "cpu",
device_count=device_count,
visible_devices=visible_devices,
total_memory_gb=total_memory_gb,
smi_command=self.smi_command,
available=available,
reason="" if available else "mthreads-gmi did not report any device",
metadata={
"torch_cuda_available": torch_cuda_available,
"torch_version": torch_version,
"supports_sharded": available,
},
)

View File

@ -0,0 +1,85 @@
# -*- coding: utf-8 -*-
"""NVIDIA CUDA accelerator adapter."""
from __future__ import annotations
from .base import (
AcceleratorInfo,
command_exists,
first_env_devices,
parse_memory_gb,
run_command,
)
class NvidiaAcceleratorAdapter:
vendor = "nvidia"
runtime = "cuda"
smi_command = "nvidia-smi"
def detect(self) -> AcceleratorInfo:
visible_devices = first_env_devices(("CUDA_VISIBLE_DEVICES",))
smi_available = command_exists(self.smi_command)
smi_count = 0
smi_memory_gb = 0.0
if smi_available:
indexes_output = run_command(
[
self.smi_command,
"--query-gpu=index",
"--format=csv,noheader,nounits",
]
)
if indexes_output:
smi_count = len([line for line in indexes_output.splitlines() if line.strip()])
memory_output = run_command(
[
self.smi_command,
"--query-gpu=memory.total",
"--format=csv,noheader",
]
)
smi_memory_gb = parse_memory_gb(memory_output)
torch_available = False
torch_count = 0
torch_memory_gb = 0.0
torch_version = ""
try:
import torch
torch_version = getattr(torch, "__version__", "")
torch_available = bool(torch.cuda.is_available())
if torch_available:
torch_count = int(torch.cuda.device_count())
torch_memory_gb = min(
torch.cuda.get_device_properties(i).total_memory / (1024**3)
for i in range(torch_count)
)
except Exception:
pass
device_count = torch_count or smi_count
if visible_devices:
device_count = min(device_count or len(visible_devices), len(visible_devices))
total_memory_gb = torch_memory_gb or smi_memory_gb
available = bool(torch_available or smi_count > 0)
reason = "" if available else "nvidia runtime not detected"
return AcceleratorInfo(
vendor=self.vendor,
runtime=self.runtime,
device="cuda:0" if available else "cpu",
device_count=device_count,
visible_devices=visible_devices,
total_memory_gb=total_memory_gb,
smi_command=self.smi_command if smi_available else None,
available=available,
reason=reason,
metadata={
"torch_cuda_available": torch_available,
"torch_version": torch_version,
"supports_sharded": available,
},
)

536
app/core/config.py 100644
View File

@ -0,0 +1,536 @@
# -*- coding: utf-8 -*-
"""
统一配置管理
ASR语音识别配置选项
"""
import os
from typing import Optional
from pathlib import Path
class Settings:
"""统一应用配置类"""
# 应用信息
APP_NAME: str = "Qwen3-ASR Server"
APP_VERSION: str = "1.0.1"
APP_DESCRIPTION: str = "Qwen3-ASR speech recognition API service"
# 服务器配置
HOST: str = "0.0.0.0"
PORT: int = 8000
DEBUG: bool = False
# 鉴权配置
API_KEY: Optional[str] = None # 从环境变量API_KEY读取,如果为None则鉴权可选
# 设备配置
ACCELERATOR: str = "auto" # auto, cpu, nvidia, metax, iluvatar, mthreads
DEVICE: str = "auto" # auto, cpu, cuda:0
ASR_DEPLOY_TOPOLOGY: str = "isolated" # isolated, sharded, auto
# 路径配置
BASE_DIR: Path = Path(__file__).parent.parent.parent
DATA_DIR: str = str(BASE_DIR / "data")
TEMP_DIR: str = str(BASE_DIR / "data" / "temp")
# 项目总模型目录。实际模型目录直接扁平化到:
# /models/{Qwen,iic,damo,...}
MODELS_DIR: str = str(BASE_DIR / "models")
# ModelScope 会在 MODELSCOPE_CACHE 下创建 models/{publisher}/{model_name}。
# 因此 cache 根目录应指向 models 的上一级,实际运行模型根目录仍由 MODELSCOPE_PATH 指定。
MODELSCOPE_CACHE: str = str(BASE_DIR)
MODELSCOPE_PATH: str = str(BASE_DIR / "models")
# 日志配置
LOG_LEVEL: str = "INFO"
LOG_FILE: Optional[str] = str(BASE_DIR / "data" / "logs" / "qwen3-asr.log")
LOG_MAX_BYTES: int = 20 * 1024 * 1024 # 20MB
LOG_BACKUP_COUNT: int = 50 # 保留50个备份文件
# ASR模型配置
WS_MAX_BUFFER_SIZE: int = 10 * 16000 # WebSocket音频缓冲区最大大小(10秒@16kHz)
FUNASR_AUTOMODEL_KWARGS = {
"trust_remote_code": False,
"disable_update": True,
"disable_pbar": True,
"disable_log": True, # 禁用FunASR的tables输出
"local_files_only": True, # 强制使用本地模型,禁止联网下载
}
ASR_MODELS_CONFIG: str = str(BASE_DIR / "app/services/asr/models.json")
ASR_ENABLE_REALTIME_PUNC: bool = True # 是否启用实时标点模型(用于中间结果展示)
VAD_MODEL: str = "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch"
PUNC_MODEL: str = "iic/punc_ct-transformer_zh-cn-common-vocab272727-pytorch"
PUNC_REALTIME_MODEL: str = (
"iic/punc_ct-transformer_zh-cn-common-vad_realtime-vocab272727"
)
# 流式ASR远场过滤配置
ASR_ENABLE_NEARFIELD_FILTER: bool = True # 是否启用远场声音过滤
ASR_NEARFIELD_RMS_THRESHOLD: float = 0.01 # RMS能量阈值(宽松模式,适合大多数场景)
# 音频处理配置
MAX_AUDIO_SIZE: int = 2048 * 1024 * 1024 # 2GB
# 批处理推理配置(GPU 真并行)
ASR_BATCH_SIZE: int = 4 # ASR 批处理大小(同时推理的片段数),建议 2-8
ASR_ENABLE_WORD_TIMESTAMPS: bool = False # 全局字词级时间戳开关;关闭时不预热 forced aligner,接口参数也默认隐藏
# 音频分段配置
MAX_SEGMENT_SEC: float = 60.0 # Max offline ASR segment duration in seconds.
# Runtime 并发配置(按 backend 独立控制)
QWEN_VLLM_SHARED_CONCURRENCY: int = 8
QWEN_VLLM_ENFORCE_EAGER: bool = True
QWEN_RUST_CPU_WORKERS: int = 4
QWEN_RUST_ASR_CONCURRENCY: int = 0
QWEN_RUST_ALIGN_CONCURRENCY: int = 0
FUNASR_WORKERS: int = 1
# 声纹数据库配置(与 Model-Test-New 使用同一套 PostgreSQL/pgvector 表结构)
SPEAKER_DB_ENABLED: bool = True
DB_USER: str = "postgres"
DB_PASSWORD: str = "postgres"
DB_NAME: str = "asr_db"
DB_HOST: str = "127.0.0.1"
DB_PORT: int = 5432
DB_POOL_MAX_SIZE: int = 5
SV_MODEL: str = "iic/speech_campplus_sv_zh-cn_16k-common"
SV_MODEL_REVISION: str = "v2.0.2"
REALTIME_SV_MODEL: str = "iic/speech_eres2netv2_sv_zh-cn_16k-common"
REALTIME_SV_MODEL_REVISION: str = ""
SV_THRESHOLD: float = 0.6
REALTIME_MAX_SEGMENT_SEC: float = 12.0
REALTIME_MAX_SEGMENT_TAIL_SEC: float = 1.6
REALTIME_FORCE_STABLE_SEGMENT_SEC: float = 8.0
REALTIME_FORCE_STABLE_MIN_CHARS: int = 24
REALTIME_MIN_PARTIAL_SEC: float = 0.45
REALTIME_PARTIAL_EMIT_INTERVAL_SEC: float = 0.25
REALTIME_PARTIAL_WINDOW_SEC: float = 8.0
REALTIME_STREAM_CHUNK_SEC: float = 1.2
REALTIME_STREAM_MAX_PENDING_CHUNKS: int = 3
REALTIME_STREAM_WINDOW_SEC: float = 8.0
REALTIME_STREAM_STABLE_TAIL_CHARS: int = 8
REALTIME_STREAM_STABLE_MIN_GROW_CHARS: int = 2
REALTIME_STREAM_DIVERGENCE_TOLERANCE_CHARS: int = 2
REALTIME_PARTIAL_HOLDBACK_CHARS: int = 6
REALTIME_LONGFORM_MIN_SEC: float = 8.0
REALTIME_LONGFORM_CHUNK_SEC: float = 6.0
REALTIME_LONGFORM_OVERLAP_SEC: float = 1.2
REALTIME_VAD_CHECK_INTERVAL_SEC: float = 0.8
REALTIME_VAD_FINALIZE_SILENCE_SEC: float = 0.6
REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC: float = 8.0
REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC: float = 12.0
REALTIME_ENABLE_DIARIZATION: bool = True
REALTIME_ENABLE_SEGMENT_REFINE: bool = False
REALTIME_DIARIZATION_MIN_SEC: float = 3.0
REALTIME_DIARIZATION_LOOKBACK_SEC: float = 3.0
REALTIME_DIARIZATION_WINDOW_SEC: float = 15.0
REALTIME_SPEAKER_MIN_SEC: float = 1.2
REALTIME_SPEAKER_CLUSTER_THRESHOLD: float = 0.75
REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD: float = 0.58
REALTIME_RECENT_UNKNOWN_SPK_THRESHOLD: float = 0.50
REALTIME_SPEAKER_CONFIRM_THRESHOLD: float = 0.62
REALTIME_REGISTRY_MIN_CLUSTER_CONFIDENCE: float = 0.72
REALTIME_SPEAKER_MAX_SLOTS: int = 8
REALTIME_SESSION_RESUME_TTL_SEC: int = 120 # WebSocket 断线后保留会话上下文的秒数;TTL 内同 session_id 可恢复
API_PREFIX: str = "/api/v1"
TASK_STATE_DIR: str = str(BASE_DIR / "data" / "tasks")
TASK_RETENTION_HOURS: int = 24
def __init__(self):
"""从环境变量读取配置"""
self._load_from_env()
self._ensure_directories()
def _load_from_env(self):
"""从环境变量加载配置"""
# 服务器配置
self.HOST = os.getenv("HOST", self.HOST)
self.PORT = int(os.getenv("PORT", str(self.PORT)))
self.DEBUG = os.getenv("DEBUG", "false").lower() == "true"
self.DATA_DIR = os.getenv("DATA_DIR", self.DATA_DIR)
self.TEMP_DIR = os.getenv("TEMP_DIR", self.TEMP_DIR)
# 日志配置
self.LOG_LEVEL = os.getenv("LOG_LEVEL", self.LOG_LEVEL)
self.LOG_FILE = os.getenv("LOG_FILE", self.LOG_FILE)
self.LOG_MAX_BYTES = int(os.getenv("LOG_MAX_BYTES", str(self.LOG_MAX_BYTES)))
self.LOG_BACKUP_COUNT = int(
os.getenv("LOG_BACKUP_COUNT", str(self.LOG_BACKUP_COUNT))
)
# 鉴权配置:空值/空白统一视为未配置
self.API_KEY = (os.getenv("API_KEY") or "").strip() or None
# 设备配置
self.ACCELERATOR = os.getenv("ACCELERATOR", self.ACCELERATOR)
self.DEVICE = os.getenv("DEVICE", self.DEVICE)
self.ASR_DEPLOY_TOPOLOGY = os.getenv(
"ASR_DEPLOY_TOPOLOGY",
self.ASR_DEPLOY_TOPOLOGY,
).strip().lower()
# 模型缓存路径
self.MODELS_DIR = os.getenv("MODELS_DIR", self.MODELS_DIR)
self.MODELSCOPE_CACHE = os.getenv("MODELSCOPE_CACHE", self.MODELSCOPE_CACHE)
self.MODELSCOPE_PATH = os.getenv("MODELSCOPE_PATH", self.MODELSCOPE_PATH)
# 给第三方库补齐默认缓存环境变量,允许用户自行覆盖
os.environ.setdefault("MODELS_DIR", self.MODELS_DIR)
os.environ.setdefault("MODELSCOPE_CACHE", self.MODELSCOPE_CACHE)
os.environ.setdefault("MODELSCOPE_PATH", self.MODELSCOPE_PATH)
# ASR模型配置
self.ASR_ENABLE_REALTIME_PUNC = (
os.getenv("ASR_ENABLE_REALTIME_PUNC", "true").lower() == "true"
)
# WebSocket缓冲区配置
self.WS_MAX_BUFFER_SIZE = int(
os.getenv("WS_MAX_BUFFER_SIZE", str(self.WS_MAX_BUFFER_SIZE))
)
# 远场过滤配置
self.ASR_ENABLE_NEARFIELD_FILTER = (
os.getenv("ASR_ENABLE_NEARFIELD_FILTER", "true").lower() == "true"
)
self.ASR_NEARFIELD_RMS_THRESHOLD = float(
os.getenv(
"ASR_NEARFIELD_RMS_THRESHOLD", str(self.ASR_NEARFIELD_RMS_THRESHOLD)
)
)
# 音频处理配置
# 支持简化格式:纯数字表示MB,或带单位(如 2048MB, 2GB)
max_audio_size_str = os.getenv("MAX_AUDIO_SIZE")
if max_audio_size_str:
self.MAX_AUDIO_SIZE = self._parse_size(max_audio_size_str)
self.ASR_BATCH_SIZE = int(
os.getenv("ASR_BATCH_SIZE", str(self.ASR_BATCH_SIZE))
)
self.ASR_ENABLE_WORD_TIMESTAMPS = (
os.getenv(
"ASR_ENABLE_WORD_TIMESTAMPS",
str(self.ASR_ENABLE_WORD_TIMESTAMPS),
).lower()
== "true"
)
self.MAX_SEGMENT_SEC = float(
os.getenv("MAX_SEGMENT_SEC", str(self.MAX_SEGMENT_SEC))
)
self.QWEN_VLLM_SHARED_CONCURRENCY = int(
os.getenv(
"QWEN_VLLM_SHARED_CONCURRENCY",
str(self.QWEN_VLLM_SHARED_CONCURRENCY),
)
)
self.QWEN_VLLM_ENFORCE_EAGER = (
os.getenv(
"QWEN_VLLM_ENFORCE_EAGER",
str(self.QWEN_VLLM_ENFORCE_EAGER),
).lower()
== "true"
)
self.QWEN_RUST_CPU_WORKERS = int(
os.getenv("QWEN_RUST_CPU_WORKERS", str(self.QWEN_RUST_CPU_WORKERS))
)
self.QWEN_RUST_ASR_CONCURRENCY = int(
os.getenv("QWEN_RUST_ASR_CONCURRENCY", str(self.QWEN_RUST_ASR_CONCURRENCY))
)
self.QWEN_RUST_ALIGN_CONCURRENCY = int(
os.getenv("QWEN_RUST_ALIGN_CONCURRENCY", str(self.QWEN_RUST_ALIGN_CONCURRENCY))
)
self.FUNASR_WORKERS = int(
os.getenv("FUNASR_WORKERS", str(self.FUNASR_WORKERS))
)
self.SPEAKER_DB_ENABLED = (
os.getenv("SPEAKER_DB_ENABLED", str(self.SPEAKER_DB_ENABLED)).lower()
== "true"
)
self.DB_USER = os.getenv("DB_USER", self.DB_USER)
self.DB_PASSWORD = os.getenv("DB_PASSWORD", self.DB_PASSWORD)
self.DB_NAME = os.getenv("DB_NAME", self.DB_NAME)
self.DB_HOST = os.getenv("DB_HOST", self.DB_HOST)
self.DB_PORT = int(os.getenv("DB_PORT", str(self.DB_PORT)))
self.DB_POOL_MAX_SIZE = int(
os.getenv("DB_POOL_MAX_SIZE", str(self.DB_POOL_MAX_SIZE))
)
self.SV_MODEL = os.getenv("SV_MODEL", self.SV_MODEL)
self.SV_MODEL_REVISION = os.getenv("SV_MODEL_REVISION", self.SV_MODEL_REVISION)
self.REALTIME_SV_MODEL = os.getenv("REALTIME_SV_MODEL", self.REALTIME_SV_MODEL)
self.REALTIME_SV_MODEL_REVISION = os.getenv(
"REALTIME_SV_MODEL_REVISION",
self.REALTIME_SV_MODEL_REVISION,
)
self.SV_THRESHOLD = float(os.getenv("SV_THRESHOLD", str(self.SV_THRESHOLD)))
self.REALTIME_MAX_SEGMENT_SEC = float(
os.getenv(
"REALTIME_MAX_SEGMENT_SEC",
str(self.REALTIME_MAX_SEGMENT_SEC),
)
)
self.REALTIME_MAX_SEGMENT_TAIL_SEC = float(
os.getenv(
"REALTIME_MAX_SEGMENT_TAIL_SEC",
str(self.REALTIME_MAX_SEGMENT_TAIL_SEC),
)
)
self.REALTIME_FORCE_STABLE_SEGMENT_SEC = float(
os.getenv(
"REALTIME_FORCE_STABLE_SEGMENT_SEC",
str(self.REALTIME_FORCE_STABLE_SEGMENT_SEC),
)
)
self.REALTIME_FORCE_STABLE_MIN_CHARS = int(
os.getenv(
"REALTIME_FORCE_STABLE_MIN_CHARS",
str(self.REALTIME_FORCE_STABLE_MIN_CHARS),
)
)
self.REALTIME_MIN_PARTIAL_SEC = float(
os.getenv(
"REALTIME_MIN_PARTIAL_SEC",
str(self.REALTIME_MIN_PARTIAL_SEC),
)
)
self.REALTIME_PARTIAL_EMIT_INTERVAL_SEC = float(
os.getenv(
"REALTIME_PARTIAL_EMIT_INTERVAL_SEC",
str(self.REALTIME_PARTIAL_EMIT_INTERVAL_SEC),
)
)
self.REALTIME_PARTIAL_WINDOW_SEC = float(
os.getenv(
"REALTIME_PARTIAL_WINDOW_SEC",
str(self.REALTIME_PARTIAL_WINDOW_SEC),
)
)
self.REALTIME_STREAM_CHUNK_SEC = float(
os.getenv(
"REALTIME_STREAM_CHUNK_SEC",
str(self.REALTIME_STREAM_CHUNK_SEC),
)
)
self.REALTIME_STREAM_MAX_PENDING_CHUNKS = int(
os.getenv(
"REALTIME_STREAM_MAX_PENDING_CHUNKS",
str(self.REALTIME_STREAM_MAX_PENDING_CHUNKS),
)
)
self.REALTIME_STREAM_WINDOW_SEC = float(
os.getenv(
"REALTIME_STREAM_WINDOW_SEC",
str(self.REALTIME_STREAM_WINDOW_SEC),
)
)
self.REALTIME_STREAM_STABLE_TAIL_CHARS = int(
os.getenv(
"REALTIME_STREAM_STABLE_TAIL_CHARS",
str(self.REALTIME_STREAM_STABLE_TAIL_CHARS),
)
)
self.REALTIME_STREAM_STABLE_MIN_GROW_CHARS = int(
os.getenv(
"REALTIME_STREAM_STABLE_MIN_GROW_CHARS",
str(self.REALTIME_STREAM_STABLE_MIN_GROW_CHARS),
)
)
self.REALTIME_STREAM_DIVERGENCE_TOLERANCE_CHARS = int(
os.getenv(
"REALTIME_STREAM_DIVERGENCE_TOLERANCE_CHARS",
str(self.REALTIME_STREAM_DIVERGENCE_TOLERANCE_CHARS),
)
)
self.REALTIME_PARTIAL_HOLDBACK_CHARS = int(
os.getenv(
"REALTIME_PARTIAL_HOLDBACK_CHARS",
str(self.REALTIME_PARTIAL_HOLDBACK_CHARS),
)
)
self.REALTIME_LONGFORM_MIN_SEC = float(
os.getenv(
"REALTIME_LONGFORM_MIN_SEC",
str(self.REALTIME_LONGFORM_MIN_SEC),
)
)
self.REALTIME_LONGFORM_CHUNK_SEC = float(
os.getenv(
"REALTIME_LONGFORM_CHUNK_SEC",
str(self.REALTIME_LONGFORM_CHUNK_SEC),
)
)
self.REALTIME_LONGFORM_OVERLAP_SEC = float(
os.getenv(
"REALTIME_LONGFORM_OVERLAP_SEC",
str(self.REALTIME_LONGFORM_OVERLAP_SEC),
)
)
self.REALTIME_VAD_CHECK_INTERVAL_SEC = float(
os.getenv(
"REALTIME_VAD_CHECK_INTERVAL_SEC",
str(self.REALTIME_VAD_CHECK_INTERVAL_SEC),
)
)
self.REALTIME_VAD_FINALIZE_SILENCE_SEC = float(
os.getenv(
"REALTIME_VAD_FINALIZE_SILENCE_SEC",
str(self.REALTIME_VAD_FINALIZE_SILENCE_SEC),
)
)
self.REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC = float(
os.getenv(
"REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC",
str(self.REALTIME_FINAL_SEGMENT_SOFT_LIMIT_SEC),
)
)
self.REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC = float(
os.getenv(
"REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC",
str(self.REALTIME_FINAL_SEGMENT_HARD_LIMIT_SEC),
)
)
self.REALTIME_ENABLE_DIARIZATION = (
os.getenv(
"REALTIME_ENABLE_DIARIZATION",
str(self.REALTIME_ENABLE_DIARIZATION),
).lower()
== "true"
)
self.REALTIME_ENABLE_SEGMENT_REFINE = (
os.getenv(
"REALTIME_ENABLE_SEGMENT_REFINE",
str(self.REALTIME_ENABLE_SEGMENT_REFINE),
).lower()
== "true"
)
self.REALTIME_DIARIZATION_MIN_SEC = float(
os.getenv(
"REALTIME_DIARIZATION_MIN_SEC",
str(self.REALTIME_DIARIZATION_MIN_SEC),
)
)
self.REALTIME_DIARIZATION_LOOKBACK_SEC = float(
os.getenv(
"REALTIME_DIARIZATION_LOOKBACK_SEC",
str(self.REALTIME_DIARIZATION_LOOKBACK_SEC),
)
)
self.REALTIME_DIARIZATION_WINDOW_SEC = float(
os.getenv(
"REALTIME_DIARIZATION_WINDOW_SEC",
str(self.REALTIME_DIARIZATION_WINDOW_SEC),
)
)
self.REALTIME_SPEAKER_MIN_SEC = float(
os.getenv(
"REALTIME_SPEAKER_MIN_SEC",
str(self.REALTIME_SPEAKER_MIN_SEC),
)
)
self.REALTIME_SPEAKER_CLUSTER_THRESHOLD = float(
os.getenv(
"REALTIME_SPEAKER_CLUSTER_THRESHOLD",
str(self.REALTIME_SPEAKER_CLUSTER_THRESHOLD),
)
)
self.REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD = float(
os.getenv(
"REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD",
str(self.REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD),
)
)
self.REALTIME_RECENT_UNKNOWN_SPK_THRESHOLD = float(
os.getenv(
"REALTIME_RECENT_UNKNOWN_SPK_THRESHOLD",
str(self.REALTIME_RECENT_UNKNOWN_SPK_THRESHOLD),
)
)
self.REALTIME_SPEAKER_CONFIRM_THRESHOLD = float(
os.getenv(
"REALTIME_SPEAKER_CONFIRM_THRESHOLD",
str(self.REALTIME_SPEAKER_CONFIRM_THRESHOLD),
)
)
self.REALTIME_REGISTRY_MIN_CLUSTER_CONFIDENCE = float(
os.getenv(
"REALTIME_REGISTRY_MIN_CLUSTER_CONFIDENCE",
str(self.REALTIME_REGISTRY_MIN_CLUSTER_CONFIDENCE),
)
)
self.REALTIME_SPEAKER_MAX_SLOTS = int(
os.getenv(
"REALTIME_SPEAKER_MAX_SLOTS",
str(self.REALTIME_SPEAKER_MAX_SLOTS),
)
)
self.REALTIME_SESSION_RESUME_TTL_SEC = int(
os.getenv(
"REALTIME_SESSION_RESUME_TTL_SEC",
str(self.REALTIME_SESSION_RESUME_TTL_SEC),
)
)
self.API_PREFIX = os.getenv("API_PREFIX", self.API_PREFIX)
self.TASK_STATE_DIR = os.getenv("TASK_STATE_DIR", self.TASK_STATE_DIR)
self.TASK_RETENTION_HOURS = int(
os.getenv("TASK_RETENTION_HOURS", str(self.TASK_RETENTION_HOURS))
)
def _parse_size(self, size_str: str) -> int:
"""解析带单位的大小字符串
支持格式:
- 纯数字:视为 MB(如 2048 = 2048MB = 2147483648 bytes)
- 带单位:如 2GB, 2048MB, 1.5GB
"""
size_str = size_str.strip().upper()
# 如果纯数字,视为 MB
if size_str.isdigit():
return int(size_str) * 1024 * 1024
# 带单位的处理
if size_str.endswith('GB'):
return int(float(size_str[:-2]) * 1024 * 1024 * 1024)
elif size_str.endswith('MB'):
return int(float(size_str[:-2]) * 1024 * 1024)
elif size_str.endswith('KB'):
return int(float(size_str[:-2]) * 1024)
else:
# 默认视为字节
return int(size_str)
def _ensure_directories(self):
"""确保必需的目录存在"""
os.makedirs(self.TEMP_DIR, exist_ok=True)
if self.LOG_FILE:
os.makedirs(os.path.dirname(self.LOG_FILE), exist_ok=True)
os.makedirs(self.MODELS_DIR, exist_ok=True)
os.makedirs(self.MODELSCOPE_CACHE, exist_ok=True)
os.makedirs(self.MODELSCOPE_PATH, exist_ok=True)
os.makedirs(self.DATA_DIR, exist_ok=True)
os.makedirs(self.TASK_STATE_DIR, exist_ok=True)
@property
def models_config_path(self) -> str:
"""获取模型配置文件的完整路径"""
return str(self.BASE_DIR / self.ASR_MODELS_CONFIG)
@property
def docs_url(self) -> Optional[str]:
"""获取文档URL"""
return "/docs"
@property
def redoc_url(self) -> Optional[str]:
"""获取ReDoc URL"""
return "/redoc"
# 全局配置实例
settings = Settings()

View File

@ -0,0 +1,222 @@
# -*- coding: utf-8 -*-
"""PostgreSQL/pgvector storage for registered speaker embeddings."""
from __future__ import annotations
import logging
from typing import Optional
import numpy as np
from app.core.config import settings
logger = logging.getLogger(__name__)
try:
import asyncpg
except ImportError: # pragma: no cover - runtime dependency is installed in Docker images.
asyncpg = None # type: ignore[assignment]
class PgSpeakerStorage:
"""Small asyncpg wrapper shared by speaker registration and ASR matching."""
def __init__(self) -> None:
self.pool: Optional[asyncpg.Pool] = None
@property
def is_enabled(self) -> bool:
return bool(settings.SPEAKER_DB_ENABLED)
@property
def is_connected(self) -> bool:
return self.pool is not None
async def _create_database_if_not_exists(self) -> None:
system_config = {
"user": settings.DB_USER,
"password": settings.DB_PASSWORD,
"database": "postgres",
"host": settings.DB_HOST,
"port": settings.DB_PORT,
}
connection: Optional[asyncpg.Connection] = None
try:
if asyncpg is None:
raise RuntimeError("asyncpg is not installed")
connection = await asyncpg.connect(**system_config)
exists = await connection.fetchval(
"SELECT 1 FROM pg_database WHERE datname = $1",
settings.DB_NAME,
)
if not exists:
await connection.execute(f'CREATE DATABASE "{settings.DB_NAME}"')
except Exception as exc:
logger.warning("尝试自动创建数据库失败,将继续连接业务库: %s", exc)
finally:
if connection is not None:
await connection.close()
async def connect(self) -> None:
if not self.is_enabled:
logger.info("声纹数据库未启用,跳过 PostgreSQL 连接")
return
if asyncpg is None:
raise RuntimeError("asyncpg is not installed")
if self.pool is not None:
return
await self._create_database_if_not_exists()
self.pool = await asyncpg.create_pool(
user=settings.DB_USER,
password=settings.DB_PASSWORD,
database=settings.DB_NAME,
host=settings.DB_HOST,
port=settings.DB_PORT,
min_size=1,
max_size=settings.DB_POOL_MAX_SIZE,
)
async with self.pool.acquire() as connection:
await connection.execute("CREATE EXTENSION IF NOT EXISTS vector;")
await connection.execute(
"""
CREATE TABLE IF NOT EXISTS speakers (
id SERIAL PRIMARY KEY,
name TEXT NOT NULL UNIQUE,
user_id TEXT,
embedding vector(192),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
"""
)
await connection.execute(
"""
DO $$
BEGIN
IF NOT EXISTS (
SELECT 1 FROM information_schema.columns
WHERE table_name='speakers' AND column_name='user_id'
) THEN
ALTER TABLE speakers ADD COLUMN user_id TEXT;
END IF;
END $$;
"""
)
await connection.execute(
"""
DO $$
BEGIN
IF NOT EXISTS (
SELECT 1
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
WHERE c.relname = 'speakers_embedding_idx'
) THEN
CREATE INDEX speakers_embedding_idx
ON speakers USING hnsw (embedding vector_cosine_ops)
WITH (m = 16, ef_construction = 64);
END IF;
END $$;
"""
)
logger.info(
"声纹数据库已连接: %s:%s/%s",
settings.DB_HOST,
settings.DB_PORT,
settings.DB_NAME,
)
async def close(self) -> None:
if self.pool is not None:
await self.pool.close()
self.pool = None
def _require_pool(self) -> asyncpg.Pool:
if self.pool is None:
raise RuntimeError("Speaker database pool is not initialized")
return self.pool
@staticmethod
def _embedding_text(embedding: np.ndarray) -> str:
return str(np.asarray(embedding, dtype=np.float32).flatten().tolist())
async def save_speaker(
self,
name: str,
embedding: np.ndarray,
user_id: Optional[str] = None,
) -> dict[str, Optional[str]]:
pool = self._require_pool()
async with pool.acquire() as connection:
row = await connection.fetchrow(
"""
INSERT INTO speakers (name, embedding, user_id)
VALUES ($1, $2, $3)
ON CONFLICT (name) DO UPDATE
SET embedding = EXCLUDED.embedding, user_id = EXCLUDED.user_id
RETURNING id, name, user_id;
""",
name,
self._embedding_text(embedding),
user_id,
)
return {
"id": str(row["id"]) if row is not None else None,
"name": row["name"] if row is not None else name,
"user_id": row["user_id"] if row is not None else user_id,
}
async def identify_speaker(
self,
embedding: np.ndarray,
threshold: float,
) -> dict[str, Optional[str]]:
pool = self._require_pool()
distance_threshold = 1.0 - float(threshold)
async with pool.acquire() as connection:
row = await connection.fetchrow(
"""
SELECT id, name, user_id, (embedding <=> $1::vector) AS distance
FROM speakers
ORDER BY embedding <=> $1::vector
LIMIT 1;
""",
self._embedding_text(embedding),
)
if row is not None and float(row["distance"]) < distance_threshold:
return {
"id": str(row["id"]),
"name": row["name"],
"user_id": row["user_id"],
}
return {"id": None, "name": None, "user_id": None}
async def list_speakers(self) -> list[dict[str, Optional[str]]]:
pool = self._require_pool()
async with pool.acquire() as connection:
rows = await connection.fetch(
"SELECT id, name, user_id, created_at FROM speakers ORDER BY id ASC"
)
return [
{
"id": str(row["id"]),
"name": row["name"],
"user_id": row["user_id"],
"created_at": row["created_at"].isoformat()
if row["created_at"]
else None,
}
for row in rows
]
async def delete_speaker(self, speaker_id: int) -> bool:
pool = self._require_pool()
async with pool.acquire() as connection:
result = await connection.execute(
"DELETE FROM speakers WHERE id = $1",
speaker_id,
)
return result != "DELETE 0"
pg_speaker_db = PgSpeakerStorage()

51
app/core/device.py 100644
View File

@ -0,0 +1,51 @@
# -*- coding: utf-8 -*-
"""Centralized device detection utility.
This module keeps the historic device helpers while delegating hardware
probing to ``app.core.accelerator``.
"""
from app.core.accelerator import get_accelerator_info
def detect_device(configured: str = "auto") -> str:
"""Resolve a device configuration string to a concrete PyTorch device.
Priority for ``"auto"``: configured accelerator > detected GPU > CPU.
Args:
configured: Value from ``settings.DEVICE`` or caller override.
Accepted: ``"auto"``, ``"cpu"``, ``"cuda:0"``, ``"npu:0"``, etc.
Returns:
A device string ready for ``torch.device()`` / FunASR / ModelScope.
"""
device = configured.strip().lower()
if device == "auto":
return get_accelerator_info().device
# Normalize bare "cuda" to "cuda:0"
if device == "cuda":
return "cuda:0"
if device == "mps":
return "cpu"
return device
def is_cuda() -> bool:
"""True when the active runtime exposes a CUDA-compatible device."""
info = get_accelerator_info()
return info.available and info.device.startswith("cuda")
def has_gpu() -> bool:
"""True when a supported accelerator is available."""
return get_accelerator_info().is_gpu
def get_vram_gb() -> float:
"""Return usable accelerator memory in GB."""
return get_accelerator_info().total_memory_gb

View File

@ -0,0 +1,224 @@
# -*- coding: utf-8 -*-
"""
统一异常处理模块
定义所有自定义异常类和错误处理函数
"""
from datetime import datetime, timezone
from fastapi import HTTPException, Request
from fastapi.exception_handlers import (
http_exception_handler as fastapi_http_exception_handler,
request_validation_exception_handler as fastapi_validation_exception_handler,
)
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
import logging
from typing import Any, Dict, Optional
logger = logging.getLogger(__name__)
def _uses_meeting_legacy_response(request: Request) -> bool:
"""Return True for Model-Test-New compatible offline meeting/speaker APIs."""
path = request.url.path.rstrip("/")
method = request.method.upper()
if path.endswith("/asr/transcriptions"):
return method == "POST"
if "/asr/transcriptions/" in path:
return method == "GET"
if path.endswith("/speakers"):
return method == "POST"
return False
def get_iso_timestamp() -> str:
"""获取ISO 8601格式的UTC时间戳"""
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
def create_error_response(
error_code: str,
message: str,
task_id: str = "",
details: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
创建标准错误响应格式
Args:
error_code: 错误代码(如 INVALID_PARAMETER)
message: 人类可读的错误信息
task_id: 任务ID(可选)
details: 额外详细信息(可选)
Returns:
标准错误响应字典
"""
return {
"error_code": error_code,
"message": message,
"task_id": task_id or "",
"timestamp": get_iso_timestamp(),
"details": details or {},
}
def get_http_status_code(status_code: int) -> int:
"""Map internal API status code to HTTP status code."""
return 500 if status_code >= 50000000 else 400
class APIException(Exception):
"""API基础异常类"""
def __init__(
self,
status_code: int,
message: str,
task_id: str = "",
error_code: str = "",
details: Optional[Dict[str, Any]] = None,
):
self.status_code = status_code
self.message = message
self.task_id = task_id
self.error_code = error_code or self._get_error_code(status_code)
self.details = details or {}
super().__init__(self.message)
def _get_error_code(self, status_code: int) -> str:
"""根据状态码获取错误代码"""
code_mapping = {
20000000: "SUCCESS",
40000000: "DEFAULT_CLIENT_ERROR",
40000001: "AUTHENTICATION_FAILED",
40000002: "INVALID_MESSAGE",
40000003: "INVALID_PARAMETER",
40000004: "IDLE_TIMEOUT",
40000005: "TOO_MANY_REQUESTS",
40000010: "TRIAL_EXPIRED",
41010101: "UNSUPPORTED_SAMPLE_RATE",
50000000: "DEFAULT_SERVER_ERROR",
50000001: "INTERNAL_GRPC_ERROR",
}
return code_mapping.get(status_code, "UNKNOWN_ERROR")
def to_dict(self) -> Dict[str, Any]:
"""
将异常转换为标准错误响应字典
Returns:
标准错误响应字典,包含 error_code, message, task_id, timestamp, details
"""
return create_error_response(
error_code=self.error_code,
message=self.message,
task_id=self.task_id,
details=self.details,
)
# 标准异常类
class AuthenticationException(APIException):
"""身份认证异常"""
def __init__(self, message: str, task_id: str = "", details: Optional[Dict[str, Any]] = None):
super().__init__(40000001, message, task_id, details=details)
class InvalidMessageException(APIException):
"""无效消息异常"""
def __init__(self, message: str, task_id: str = "", details: Optional[Dict[str, Any]] = None):
super().__init__(40000002, message, task_id, details=details)
class InvalidParameterException(APIException):
"""无效参数异常"""
def __init__(self, message: str, task_id: str = "", details: Optional[Dict[str, Any]] = None):
super().__init__(40000003, message, task_id, details=details)
class UnsupportedSampleRateException(APIException):
"""不支持的采样率异常"""
def __init__(self, message: str, task_id: str = "", details: Optional[Dict[str, Any]] = None):
super().__init__(41010101, message, task_id, details=details)
class DefaultServerErrorException(APIException):
"""默认服务端错误异常"""
def __init__(self, message: str, task_id: str = "", details: Optional[Dict[str, Any]] = None):
super().__init__(50000000, message, task_id, details=details)
# 异常处理器
async def api_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""API异常处理器"""
# FastAPI 只会在抛出 APIException 时调用此处理器
api_exc = exc if isinstance(exc, APIException) else APIException(50000000, str(exc))
logger.error(f"[{api_exc.task_id}] API异常: {api_exc.message}")
# 使用标准错误格式
response_data = api_exc.to_dict()
# 确定HTTP状态码
http_status_code = get_http_status_code(api_exc.status_code)
return JSONResponse(
content=response_data,
headers={"task_id": api_exc.task_id} if api_exc.task_id else {},
status_code=http_status_code,
)
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""FastAPI HTTPException handler for legacy meeting/speaker endpoints."""
http_exc = exc if isinstance(exc, HTTPException) else HTTPException(500, str(exc))
if not _uses_meeting_legacy_response(request):
return await fastapi_http_exception_handler(request, http_exc)
detail = http_exc.detail
if isinstance(detail, dict):
return JSONResponse(
status_code=http_exc.status_code,
content={
"code": detail.get("code", 4001),
"message": detail.get("message", "Request failed"),
},
)
return JSONResponse(
status_code=http_exc.status_code,
content={"code": 4001, "message": str(detail)},
)
async def validation_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""FastAPI/Pydantic validation handler for legacy meeting/speaker endpoints."""
validation_exc = exc if isinstance(exc, RequestValidationError) else None
if validation_exc is None:
return JSONResponse(
status_code=422,
content={"code": 1001, "message": "Invalid request"},
)
if not _uses_meeting_legacy_response(request):
return await fastapi_validation_exception_handler(request, validation_exc)
errors = validation_exc.errors()
message = str(errors[0]["msg"]) if errors else "Invalid request"
return JSONResponse(status_code=422, content={"code": 1001, "message": message})
async def general_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""通用异常处理器"""
logger.error(f"未处理的异常: {str(exc)}", exc_info=True)
# 使用标准错误格式
response_data = create_error_response(
error_code="DEFAULT_SERVER_ERROR",
message=f"内部服务错误: {str(exc)}",
)
return JSONResponse(content=response_data, status_code=500)

View File

@ -0,0 +1,185 @@
# -*- coding: utf-8 -*-
"""
异步执行器模块
用于将同步的模型推理调用放入线程池执行,避免阻塞事件循环,
实现真正的多路并发处理。
设计要点:
1. 使用 ThreadPoolExecutor 而非 ProcessPoolExecutor
- 模型已加载在内存中,进程间无法共享
- GPU操作会自动释放GIL,线程池足以实现并发
2. 对于流式生成器,使用 asyncio.Queue 实现异步迭代
3. 线程池大小根据使用场景配置:
- CPU推理:受GIL限制,多线程并发收益有限,但可以让I/O不阻塞
- GPU推理:CUDA操作释放GIL,可以实现真正并发
"""
import os
import asyncio
import logging
from concurrent.futures import ThreadPoolExecutor
from typing import Callable, TypeVar, Generator, AsyncGenerator, Optional
from functools import partial
logger = logging.getLogger(__name__)
# 类型变量
T = TypeVar("T")
# 全局线程池执行器
# 默认线程数:max(4, CPU核心数),可通过环境变量覆盖
_DEFAULT_WORKERS = max(4, os.cpu_count() or 4)
_MAX_WORKERS = int(os.getenv("INFERENCE_THREAD_POOL_SIZE", str(_DEFAULT_WORKERS)))
_executor: Optional[ThreadPoolExecutor] = None
def get_executor() -> ThreadPoolExecutor:
"""获取全局线程池执行器(懒加载)"""
global _executor
if _executor is None:
_executor = ThreadPoolExecutor(
max_workers=_MAX_WORKERS,
thread_name_prefix="inference_worker"
)
logger.info(f"推理线程池已创建,最大工作线程数: {_MAX_WORKERS}")
return _executor
def shutdown_executor():
"""关闭线程池执行器"""
global _executor
if _executor is not None:
_executor.shutdown(wait=True)
_executor = None
logger.info("推理线程池已关闭")
async def run_sync(func: Callable[..., T], *args, **kwargs) -> T:
"""
在线程池中执行同步函数,不阻塞事件循环
Args:
func: 同步函数
*args: 位置参数
**kwargs: 关键字参数
Returns:
函数返回值
Example:
result = await run_sync(model.generate, input=audio_array, cache=cache)
"""
loop = asyncio.get_running_loop()
executor = get_executor()
# 使用 partial 绑定参数
if kwargs:
func_with_args = partial(func, *args, **kwargs)
else:
func_with_args = partial(func, *args) if args else func
return await loop.run_in_executor(executor, func_with_args)
async def run_sync_generator(
generator_func: Callable[..., Generator[T, None, None]],
*args,
**kwargs
) -> AsyncGenerator[T, None]:
"""
将同步生成器转换为异步生成器,在线程池中执行
用于流式处理等需要逐步产出结果的场景。
Args:
generator_func: 返回生成器的同步函数
*args: 位置参数
**kwargs: 关键字参数
Yields:
生成器的每个产出值
Example:
async for chunk in run_sync_generator(model.inference_sft, text, voice, stream=True):
await websocket.send_bytes(chunk)
"""
loop = asyncio.get_running_loop()
executor = get_executor()
queue: asyncio.Queue = asyncio.Queue()
# 标记生成器结束的哨兵值
_SENTINEL = object()
def producer():
"""在线程中运行生成器,将结果放入队列"""
try:
gen = generator_func(*args, **kwargs)
for item in gen:
# 使用 call_soon_threadsafe 安全地将结果放入队列
loop.call_soon_threadsafe(queue.put_nowait, item)
except BaseException as e:
# 发生异常时,记录日志并将异常放入队列
logger.error(f"生成器执行异常: {type(e).__name__}: {e}")
loop.call_soon_threadsafe(queue.put_nowait, e)
finally:
# 发送结束标记
loop.call_soon_threadsafe(queue.put_nowait, _SENTINEL)
# 在线程池中启动生产者
future = executor.submit(producer)
try:
while True:
item = await queue.get()
if item is _SENTINEL:
break
# 使用 BaseException 捕获所有异常类型(包括 KeyboardInterrupt 等)
if isinstance(item, BaseException):
raise item
yield item
finally:
# 确保检查线程是否有未捕获的异常
if future.done():
try:
# 如果线程已完成,检查是否有异常
future.result()
except Exception as e:
logger.error(f"生成器线程异常: {type(e).__name__}: {e}")
elif not future.cancelled():
future.cancel()
class AsyncInferenceWrapper:
"""
异步推理包装器
将同步的模型推理方法包装为异步方法,方便复用。
Example:
wrapper = AsyncInferenceWrapper(asr_engine.realtime_model)
result = await wrapper.generate(input=audio_array, cache=cache)
"""
def __init__(self, model):
self._model = model
async def generate(self, *args, **kwargs):
"""异步调用模型的 generate 方法"""
return await run_sync(self._model.generate, *args, **kwargs)
async def inference_sft(self, *args, **kwargs):
"""异步流式调用模型的 inference_sft 方法"""
async for item in run_sync_generator(self._model.inference_sft, *args, **kwargs):
yield item
async def inference_zero_shot(self, *args, **kwargs):
"""异步流式调用模型的 inference_zero_shot 方法"""
async for item in run_sync_generator(self._model.inference_zero_shot, *args, **kwargs):
yield item

View File

@ -0,0 +1,216 @@
# -*- coding: utf-8 -*-
"""Rule-based hotword correction for ASR post-processing."""
from __future__ import annotations
import re
from difflib import SequenceMatcher
from typing import Any
try:
from pypinyin import lazy_pinyin
except Exception: # pragma: no cover - optional runtime enhancement.
lazy_pinyin = None # type: ignore[assignment]
_HOTWORD_PROMPT_LEAK_PATTERNS = (
re.compile(r"^\s*Use this context when resolving named entities:\s*", re.IGNORECASE),
re.compile(r"^\s*(?:上下文信息[::]?\s*)?热词列表[::]\s*[[\[][^\]]]{0,500}[]\]]\s*[,。,::;;\s]*"),
)
def parse_hotwords(hotwords: str | list[dict[str, object]] | None) -> list[dict[str, Any]]:
if isinstance(hotwords, list):
return _extract_hotword_entries(hotwords)
raw_text = str(hotwords or "").strip()
if not raw_text:
return []
tokens = [part for part in re.split(r"[\s,,;;]+", raw_text) if part]
entries: list[dict[str, Any]] = []
index = 0
order = 0
while index < len(tokens):
word = tokens[index].strip()
if not word:
index += 1
continue
weight = 1.0
if index + 1 < len(tokens):
try:
weight = float(tokens[index + 1])
index += 2
except ValueError:
index += 1
else:
index += 1
entries.append({"text": word, "weight": weight, "order": order})
order += 1
return entries
def format_hotword_prompt_context(hotwords: str | list[dict[str, object]] | None) -> str:
entries = parse_hotwords(hotwords)
if not entries:
return ""
unique_words: list[str] = []
for entry in entries:
word = str(entry["text"]).strip()
if word and word not in unique_words:
unique_words.append(word)
if not unique_words:
return ""
return f"热词列表:[{', '.join(unique_words)}]"
def strip_hotword_prompt_leakage(text: str) -> str:
cleaned = str(text or "")
changed = True
while changed and cleaned:
changed = False
for pattern in _HOTWORD_PROMPT_LEAK_PATTERNS:
updated, count = pattern.subn("", cleaned, count=1)
if count:
cleaned = updated
changed = True
return cleaned.strip()
def _extract_hotword_entries(hotwords: list[dict[str, object]] | None) -> list[dict[str, Any]]:
entries: list[dict[str, Any]] = []
for index, item in enumerate(hotwords or []):
hotword_text = str(item.get("hotword", "")).strip()
if not hotword_text:
continue
try:
weight = float(item.get("weight", 1.0))
except Exception:
weight = 1.0
entries.append({"text": hotword_text, "weight": weight, "order": index})
return entries
def _normalize_ascii_token(text: str) -> str:
return re.sub(r"[^a-z0-9]+", "", str(text or "").lower())
def _common_prefix_length(left: str, right: str) -> int:
matched = 0
for left_char, right_char in zip(left, right):
if left_char != right_char:
break
matched += 1
return matched
def _common_suffix_length(left: str, right: str) -> int:
return _common_prefix_length(left[::-1], right[::-1])
def _normalize_cjk_pinyin(text: str) -> tuple[str, ...]:
if lazy_pinyin is None:
return ()
return tuple(part.strip().lower() for part in lazy_pinyin(str(text or ""), errors="ignore") if str(part).strip())
def _replace_ascii_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
updated_text = str(text or "")
changed = False
matched_hotwords: list[str] = []
ascii_pattern = re.compile(r"[A-Za-z][A-Za-z0-9\s._-]{0,40}")
for entry in hotword_entries:
hotword = str(entry["text"])
if not re.search(r"[A-Za-z]", hotword):
continue
normalized_hotword = _normalize_ascii_token(hotword)
if not normalized_hotword:
continue
def _replace_match(match: re.Match[str]) -> str:
nonlocal changed
candidate = match.group(0)
if _normalize_ascii_token(candidate) == normalized_hotword and candidate != hotword:
changed = True
if hotword not in matched_hotwords:
matched_hotwords.append(hotword)
return hotword
return candidate
updated_text = ascii_pattern.sub(_replace_match, updated_text)
return updated_text, changed, matched_hotwords
def _score_cjk_candidate(candidate: str, hotword_entry: dict[str, Any]) -> tuple[int, float, int, float, int] | None:
hotword = str(hotword_entry["text"])
if candidate == hotword or len(candidate) != len(hotword):
return None
candidate_pinyin = _normalize_cjk_pinyin(candidate)
hotword_pinyin = _normalize_cjk_pinyin(hotword)
if candidate_pinyin and candidate_pinyin == hotword_pinyin:
return (3, float(hotword_entry["weight"]), len(hotword), 1.0, -int(hotword_entry["order"]))
if len(hotword) <= 2:
return None
prefix_length = _common_prefix_length(candidate, hotword)
suffix_length = _common_suffix_length(candidate, hotword)
similarity = SequenceMatcher(None, candidate, hotword).ratio()
if prefix_length >= len(hotword) - 1 or suffix_length >= len(hotword) - 1:
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
if similarity >= 0.67 and prefix_length >= 1 and suffix_length >= 1:
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
return None
def _replace_cjk_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
updated_text = str(text or "")
changed = False
matched_hotwords: list[str] = []
chinese_entries = [entry for entry in hotword_entries if re.fullmatch(r"[\u4e00-\u9fff]{2,12}", str(entry["text"]))]
if not chinese_entries:
return updated_text, False, matched_hotwords
candidate_lengths = sorted({len(str(entry["text"])) for entry in chinese_entries})
token_pattern = re.compile(r"[\u4e00-\u9fff]{2,24}")
def _replace_match(match: re.Match[str]) -> str:
nonlocal changed
token = match.group(0)
best_choice: tuple[tuple[int, float, int, float, int], int, int, str] | None = None
for hotword_length in candidate_lengths:
if hotword_length > len(token):
continue
scoped_entries = [entry for entry in chinese_entries if len(str(entry["text"])) == hotword_length]
for index in range(0, len(token) - hotword_length + 1):
candidate = token[index:index + hotword_length]
for entry in scoped_entries:
score = _score_cjk_candidate(candidate, entry)
if score is None:
continue
current_choice = (score, index, hotword_length, str(entry["text"]))
if best_choice is None or current_choice > best_choice:
best_choice = current_choice
if best_choice is None:
return token
_, begin_index, hotword_length, hotword = best_choice
token_chars = list(token)
token_chars[begin_index:begin_index + hotword_length] = list(hotword)
changed = True
if hotword not in matched_hotwords:
matched_hotwords.append(hotword)
return "".join(token_chars)
updated_text = token_pattern.sub(_replace_match, updated_text)
return updated_text, changed, matched_hotwords
def apply_hotword_rules(text: str, hotwords: str | list[dict[str, object]] | None = None) -> tuple[str, bool, list[str]]:
hotword_entries = parse_hotwords(hotwords)
if not hotword_entries:
return str(text or ""), False, []
updated_text, ascii_changed, ascii_matches = _replace_ascii_hotwords(str(text or ""), hotword_entries)
updated_text, cjk_changed, cjk_matches = _replace_cjk_hotwords(updated_text, hotword_entries)
matched_hotwords: list[str] = []
for hotword_text in ascii_matches + cjk_matches:
if hotword_text not in matched_hotwords:
matched_hotwords.append(hotword_text)
return updated_text, ascii_changed or cjk_changed, matched_hotwords

342
app/core/logging.py 100644
View File

@ -0,0 +1,342 @@
# -*- coding: utf-8 -*-
"""
日志配置模块
统一的日志配置和管理,支持多 Worker 模式
"""
import logging
import logging.handlers
import sys
import os
import json
from datetime import datetime, timezone
from typing import Optional, Dict, Any
from pathlib import Path
from .config import settings
class StructuredLogFormatter(logging.Formatter):
"""结构化JSON日志格式化器
将日志记录格式化为JSON格式,支持extra字段传递结构化数据。
输出示例:
{
"timestamp": "2025-01-31T12:00:00Z",
"level": "INFO",
"logger": "app.services.asr",
"message": "推理完成",
"task_id": "xxx",
"duration_ms": 1234,
"audio_duration_sec": 60,
"rtf": 0.02,
"model_id": "qwen3-asr-1.7b"
}
"""
def __init__(self, include_extra: bool = True):
"""
Args:
include_extra: 是否包含extra字段中的结构化数据
"""
super().__init__()
self.include_extra = include_extra
self._reserved_attrs = {
'name', 'msg', 'args', 'levelname', 'levelno', 'pathname',
'filename', 'module', 'exc_info', 'exc_text', 'stack_info',
'lineno', 'funcName', 'created', 'msecs', 'relativeCreated',
'thread', 'threadName', 'processName', 'process', 'getMessage',
'message', 'asctime'
}
def format(self, record: logging.LogRecord) -> str:
"""将日志记录格式化为JSON"""
log_data: Dict[str, Any] = {
"timestamp": datetime.fromtimestamp(record.created, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z",
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
}
# 添加worker_id(多worker模式下)
workers = int(os.getenv("WORKERS", "1"))
if workers > 1:
log_data["worker_id"] = f"worker-{os.getpid()}"
# 添加异常信息(如果有)
if record.exc_info:
exc_type = record.exc_info[0]
exc_value = record.exc_info[1]
if exc_type and exc_value:
log_data["exception"] = {
"type": exc_type.__name__,
"message": str(exc_value)
}
# 添加extra字段中的结构化数据
if self.include_extra:
extra_data = self._extract_extra_data(record)
if extra_data:
log_data.update(extra_data)
return json.dumps(log_data, ensure_ascii=False, default=str)
def _extract_extra_data(self, record: logging.LogRecord) -> Dict[str, Any]:
"""从日志记录中提取extra数据"""
extra_data = {}
for key, value in record.__dict__.items():
if key not in self._reserved_attrs and not key.startswith('_'):
extra_data[key] = value
return extra_data
class HybridLogFormatter(logging.Formatter):
"""混合日志格式化器
根据日志内容自动选择格式:
- 包含结构化数据(extra字段)的日志使用JSON格式
- 普通日志使用文本格式
这允许在代码中逐步迁移到结构化日志,同时保持可读性。
"""
def __init__(
self,
text_format: Optional[str] = None,
json_formatter: Optional[StructuredLogFormatter] = None,
):
super().__init__()
self.text_format = text_format or "%(asctime)s - %(name)s - %(levelname)s - %(message)s"
self.json_formatter = json_formatter or StructuredLogFormatter()
self._reserved_attrs = {
'name', 'msg', 'args', 'levelname', 'levelno', 'pathname',
'filename', 'module', 'exc_info', 'exc_text', 'stack_info',
'lineno', 'funcName', 'created', 'msecs', 'relativeCreated',
'thread', 'threadName', 'processName', 'process', 'getMessage',
'message', 'asctime'
}
def format(self, record: logging.LogRecord) -> str:
"""根据内容选择格式"""
# 检查是否有extra数据
has_extra = self._has_extra_data(record)
if has_extra:
return self.json_formatter.format(record)
else:
# 使用文本格式
text_formatter = logging.Formatter(self.text_format)
return text_formatter.format(record)
def _has_extra_data(self, record: logging.LogRecord) -> bool:
"""检查日志记录是否包含extra数据"""
for key in record.__dict__.keys():
if key not in self._reserved_attrs and not key.startswith('_'):
return True
return False
def get_structured_logger(name: str) -> logging.Logger:
"""获取支持结构化日志的记录器
这是一个便捷函数,返回一个配置好的日志记录器,
可以直接使用 extra 参数记录结构化数据。
示例:
logger = get_structured_logger(__name__)
logger.info(
"推理完成",
extra={
"duration_ms": 1234,
"audio_duration_sec": 60,
"rtf": 0.02,
"model_id": "qwen3-asr-1.7b"
}
)
Args:
name: 记录器名称
Returns:
配置好的日志记录器
"""
return logging.getLogger(name)
def log_inference_metrics(
logger: logging.Logger,
message: str,
task_id: Optional[str] = None,
duration_ms: Optional[float] = None,
audio_duration_sec: Optional[float] = None,
model_id: Optional[str] = None,
status: str = "success",
**kwargs
) -> None:
"""记录推理性能指标
这是一个辅助函数,用于统一记录推理性能指标。
Args:
logger: 日志记录器
message: 日志消息
task_id: 任务ID
duration_ms: 推理耗时(毫秒)
audio_duration_sec: 音频时长(秒)
model_id: 模型ID
status: 状态(success/error)
**kwargs: 其他结构化数据
"""
extra: Dict[str, Any] = {
"status": status,
}
if task_id:
extra["task_id"] = task_id
if duration_ms is not None:
extra["duration_ms"] = round(duration_ms, 2)
if audio_duration_sec is not None:
extra["audio_duration_sec"] = round(audio_duration_sec, 2)
if model_id:
extra["model_id"] = model_id
# 计算RTF(实时率)
if duration_ms is not None and audio_duration_sec is not None and audio_duration_sec > 0:
rtf = (duration_ms / 1000) / audio_duration_sec
extra["rtf"] = round(rtf, 4)
# 添加其他数据
extra.update(kwargs)
logger.info(message, extra=extra)
def get_worker_id() -> str:
"""获取当前 Worker ID
Returns:
Worker 标识符,格式为 'worker-{pid}' 或 'main'
"""
# 检查是否在多 worker 模式下
workers = int(os.getenv("WORKERS", "1"))
if workers > 1:
return f"worker-{os.getpid()}"
return "main"
def setup_logging(
level: Optional[str] = None,
log_file: Optional[str] = None,
format_string: Optional[str] = None,
max_bytes: Optional[int] = None,
backup_count: Optional[int] = None,
worker_id: Optional[str] = None,
use_structured: bool = False,
) -> None:
"""设置应用日志配置
Args:
level: 日志级别
log_file: 日志文件路径
format_string: 日志格式字符串
max_bytes: 单个日志文件最大大小(字节)
backup_count: 保留的备份文件数量
worker_id: Worker 标识符(多 Worker 模式下使用)
use_structured: 是否使用结构化JSON日志格式
"""
# 使用传入的参数或配置文件中的设置
log_level = level or settings.LOG_LEVEL
log_file_path = log_file or settings.LOG_FILE
max_file_size = max_bytes or settings.LOG_MAX_BYTES
backup_files = backup_count or settings.LOG_BACKUP_COUNT
# 获取 Worker ID
current_worker_id = worker_id or get_worker_id()
workers = int(os.getenv("WORKERS", "1"))
worker_log_path: Optional[Path] = None
# 确定日志格式
if use_structured:
# 使用结构化日志格式
formatter: logging.Formatter = StructuredLogFormatter()
else:
# 使用混合格式(普通日志文本,带extra的JSON)
if workers > 1:
text_format = format_string or f"%(asctime)s - [{current_worker_id}] - %(name)s - %(levelname)s - %(message)s"
else:
text_format = format_string or "%(asctime)s - %(name)s - %(levelname)s - %(message)s"
formatter = HybridLogFormatter(text_format=text_format)
# 创建处理器列表
stream_handler = logging.StreamHandler(sys.stdout)
stream_handler.setFormatter(formatter)
handlers: list[logging.Handler] = [stream_handler]
# 确定日志文件路径
if log_file_path:
log_path = Path(log_file_path)
else:
log_path = Path("logs/qwen3-asr.log")
# 确保日志目录存在
log_dir = log_path.parent
log_dir.mkdir(parents=True, exist_ok=True)
# 多 Worker 模式下,每个 Worker 使用独立的日志文件
if workers > 1:
# Example: qwen3-asr.log -> qwen3-asr.worker-12345.log
worker_log_path = log_dir / f"{log_path.stem}.{current_worker_id}{log_path.suffix}"
# Worker 专属日志文件
worker_file_handler = logging.handlers.RotatingFileHandler(
worker_log_path,
maxBytes=max_file_size,
backupCount=backup_files,
encoding="utf-8",
)
worker_file_handler.setFormatter(formatter)
handlers.append(worker_file_handler)
# 同时也写入主日志文件(汇总所有 Worker 的日志)
main_file_handler = logging.handlers.RotatingFileHandler(
log_path,
maxBytes=max_file_size,
backupCount=backup_files,
encoding="utf-8",
)
main_file_handler.setFormatter(formatter)
handlers.append(main_file_handler)
else:
# 单 Worker 模式,只写入主日志文件
file_handler = logging.handlers.RotatingFileHandler(
log_path,
maxBytes=max_file_size,
backupCount=backup_files,
encoding="utf-8",
)
file_handler.setFormatter(formatter)
handlers.append(file_handler)
# 配置根日志记录器
logging.basicConfig(
level=getattr(logging, log_level.upper()),
handlers=handlers,
force=True, # 强制重新配置
)
# 设置第三方库的日志级别(由LOG_LEVEL控制)
third_party_level = getattr(logging, log_level.upper())
logging.getLogger("urllib3").setLevel(third_party_level)
logging.getLogger("requests").setLevel(third_party_level)
logging.getLogger("httpx").setLevel(third_party_level)
logging.getLogger("httpcore").setLevel(third_party_level)
# 始终禁用噪音特别大的库
logging.getLogger("numba").setLevel(logging.WARNING)
logging.getLogger("numba.core").setLevel(logging.WARNING)
logging.getLogger("numba.core.ssa").setLevel(logging.WARNING)
# 多 Worker 模式下记录启动日志
if workers > 1 and worker_log_path:
logger = logging.getLogger(__name__)
logger.info(f"Worker {current_worker_id} 日志系统已初始化,日志文件: {worker_log_path}")

View File

@ -0,0 +1,173 @@
# -*- coding: utf-8 -*-
"""
安全相关功能
包含鉴权、token验证等安全功能
"""
from typing import Optional
from fastapi import Request
from .config import settings
TOKEN_HEADER_NAME = "X-NLS-Token"
AUTH_OPTIONAL_PLACEHOLDER = "optional"
WEBSOCKET_QUERY_TOKEN_KEYS = ("token", "x_nls_token", "X-NLS-Token")
def normalize_token(token: Optional[str]) -> Optional[str]:
"""将 token 归一化为非空字符串或 None。"""
if token is None:
return None
normalized = token.strip()
return normalized or None
def get_expected_api_key(expected_token: Optional[str] = None) -> Optional[str]:
"""获取归一化后的期望 API_KEY。"""
if expected_token is not None:
return normalize_token(expected_token)
return normalize_token(settings.API_KEY)
def mask_sensitive_data(
data: str, mask_char: str = "*", keep_prefix: int = 4, keep_suffix: int = 4
) -> str:
"""遮盖敏感数据
Args:
data: 需要遮盖的数据
mask_char: 遮盖字符
keep_prefix: 保留前缀字符数
keep_suffix: 保留后缀字符数
Returns:
遮盖后的数据
"""
if not data or len(data) <= keep_prefix + keep_suffix:
return data
prefix = data[:keep_prefix]
suffix = data[-keep_suffix:] if keep_suffix > 0 else ""
mask_length = len(data) - keep_prefix - keep_suffix
mask = mask_char * mask_length
return f"{prefix}{mask}{suffix}"
def validate_token_value(token: Optional[str], expected_token: Optional[str] = None) -> bool:
"""验证访问令牌
Args:
token: 客户端提供的token
expected_token: 期望的token值(从环境变量读取),如果为None则鉴权可选
Returns:
bool: 验证结果
"""
normalized_expected_token = get_expected_api_key(expected_token)
if not normalized_expected_token:
return True
normalized_token = normalize_token(token)
if not normalized_token:
return False
# 简单的token格式验证(长度检查)
if len(normalized_token) < 10:
return False
# 验证token是否匹配
if normalized_token != normalized_expected_token:
return False
return True
def extract_header_token(request: Request) -> Optional[str]:
"""从标准头部提取 token。"""
return normalize_token(request.headers.get(TOKEN_HEADER_NAME))
def extract_bearer_token(request: Request) -> Optional[str]:
"""从 Authorization: Bearer 提取 token。"""
auth_header = request.headers.get("Authorization")
if not auth_header:
return None
scheme, _, value = auth_header.partition(" ")
if scheme.lower() != "bearer":
return None
return normalize_token(value)
def extract_openai_token(request: Request) -> Optional[str]:
"""OpenAI 兼容接口鉴权:优先 Bearer,其次 X-NLS-Token。"""
return extract_bearer_token(request) or extract_header_token(request)
def extract_websocket_token(websocket) -> Optional[str]:
"""从 WebSocket 连接中提取 token。"""
if hasattr(websocket, "headers"):
token = normalize_token(websocket.headers.get(TOKEN_HEADER_NAME))
if token:
return token
if hasattr(websocket, "query_params"):
for key in WEBSOCKET_QUERY_TOKEN_KEYS:
token = normalize_token(websocket.query_params.get(key))
if token:
return token
return None
def _validate_resolved_token(
token: Optional[str],
missing_message: str,
expected_token: Optional[str] = None,
) -> tuple[bool, str]:
"""统一 token 校验逻辑。"""
expected = get_expected_api_key(expected_token)
normalized_token = normalize_token(token)
if not expected:
return True, normalized_token or AUTH_OPTIONAL_PLACEHOLDER
if not normalized_token:
return False, missing_message
if not validate_token_value(normalized_token, expected):
masked_token = mask_sensitive_data(normalized_token)
return False, f"Gateway:ACCESS_DENIED:The token '{masked_token}' is invalid!"
return True, normalized_token
def validate_token(request: Request, task_id: str = "") -> tuple[bool, str]:
"""验证X-NLS-Token头部"""
_ = task_id
token = extract_header_token(request)
return _validate_resolved_token(token, "缺少X-NLS-Token头部")
def validate_openai_token(request: Request, task_id: str = "") -> tuple[bool, str]:
"""验证 OpenAI 兼容接口 token(Bearer/X-NLS-Token)。"""
_ = task_id
token = extract_openai_token(request)
return _validate_resolved_token(token, "缺少Authorization Bearer或X-NLS-Token头部")
def validate_token_websocket(token: str, task_id: str = "") -> tuple[bool, str]:
"""验证WebSocket连接中的token"""
_ = task_id
return _validate_resolved_token(token, "缺少token参数")
def validate_websocket_token(websocket, task_id: str = "") -> tuple[bool, str]:
"""验证 WebSocket 连接 token(header/query 参数)。"""
_ = task_id
token = extract_websocket_token(websocket)
return _validate_resolved_token(
token,
"缺少鉴权信息,请通过 X-NLS-Token header 或 token/x_nls_token 查询参数传入",
)

View File

@ -0,0 +1,280 @@
# -*- coding: utf-8 -*-
"""Persistent task state store for meeting-style offline jobs."""
from __future__ import annotations
import json
import time
from pathlib import Path
from threading import RLock
from typing import Any, Optional
from app.core.config import settings
tasks_db: dict[str, dict[str, Any]] = {}
_tasks_lock = RLock()
_last_cleanup_at = 0
def _task_state_dir() -> Path:
return Path(settings.TASK_STATE_DIR)
def _task_file_path(task_id: str) -> Path:
return _task_state_dir() / f"{task_id}.json"
def _task_result_file_path(task_id: str) -> Path:
return _task_state_dir() / f"{task_id}.result.json"
def _retention_seconds() -> int:
return max(0, int(settings.TASK_RETENTION_HOURS)) * 3600
def _is_task_expired(task_payload: dict[str, Any], now_ts: Optional[int] = None) -> bool:
retention_seconds = _retention_seconds()
if retention_seconds <= 0:
return False
current_time = int(now_ts or time.time())
base_timestamp = int(task_payload.get("updated_at") or task_payload.get("created_at") or 0)
if base_timestamp <= 0:
return False
return current_time - base_timestamp >= retention_seconds
def _write_json_file(file_path: Path, payload: dict[str, Any]) -> None:
file_path.parent.mkdir(parents=True, exist_ok=True)
temporary_file_path = file_path.with_suffix(f"{file_path.suffix}.tmp")
temporary_file_path.write_text(
json.dumps(payload, ensure_ascii=False),
encoding="utf-8",
)
temporary_file_path.replace(file_path)
def _write_task_file(task_id: str, payload: dict[str, Any]) -> None:
_write_json_file(_task_file_path(task_id), payload)
def _remove_task_file(task_id: str) -> None:
task_file_path = _task_file_path(task_id)
if task_file_path.exists():
task_file_path.unlink()
def _write_task_result_file(task_id: str, payload: dict[str, Any]) -> None:
_write_json_file(_task_result_file_path(task_id), payload)
def _remove_task_result_file(task_id: str) -> None:
task_result_file_path = _task_result_file_path(task_id)
if task_result_file_path.exists():
task_result_file_path.unlink()
def _load_task_payload_from_file(task_file_path: Path) -> Optional[tuple[str, dict[str, Any]]]:
try:
payload = json.loads(task_file_path.read_text(encoding="utf-8"))
except Exception:
return None
if not isinstance(payload, dict):
return None
task_id = str(payload.get("task_id") or task_file_path.stem).strip()
if not task_id:
return None
payload["task_id"] = task_id
return task_id, payload
def _load_task_result_payload(task_id: str) -> Optional[dict[str, Any]]:
task_result_file_path = _task_result_file_path(task_id)
if not task_result_file_path.exists():
return None
try:
payload = json.loads(task_result_file_path.read_text(encoding="utf-8"))
except Exception:
return None
if not isinstance(payload, dict):
return None
return payload
def _split_task_payload(task_id: str, payload: dict[str, Any], *, persist_legacy_result: bool = False) -> tuple[dict[str, Any], bool]:
status_payload = dict(payload)
has_result = "result" in status_payload
result_payload = status_payload.pop("result", None)
if has_result:
if isinstance(result_payload, dict):
_write_task_result_file(task_id, result_payload)
elif result_payload is None:
_remove_task_result_file(task_id)
if persist_legacy_result:
_write_task_file(task_id, status_payload)
return status_payload, has_result
def _status_file_paths() -> list[Path]:
return [
task_file_path
for task_file_path in _task_state_dir().glob("*.json")
if not task_file_path.name.endswith(".result.json")
]
def load_tasks_from_disk() -> None:
with _tasks_lock:
tasks_db.clear()
task_directory = _task_state_dir()
task_directory.mkdir(parents=True, exist_ok=True)
current_time = int(time.time())
for task_file_path in _status_file_paths():
loaded_item = _load_task_payload_from_file(task_file_path)
if loaded_item is None:
try:
task_file_path.unlink()
except Exception:
pass
continue
task_id, payload = loaded_item
payload, _ = _split_task_payload(task_id, payload, persist_legacy_result=True)
if _is_task_expired(payload, now_ts=current_time):
try:
task_file_path.unlink()
except Exception:
pass
try:
_remove_task_result_file(task_id)
except Exception:
pass
continue
tasks_db[task_id] = payload
def cleanup_expired_tasks(force: bool = False) -> None:
global _last_cleanup_at
current_time = int(time.time())
if not force and current_time - _last_cleanup_at < 60:
return
with _tasks_lock:
expired_task_ids = [
task_id
for task_id, task_payload in tasks_db.items()
if _is_task_expired(task_payload, now_ts=current_time)
]
for task_id in expired_task_ids:
tasks_db.pop(task_id, None)
try:
_remove_task_file(task_id)
except Exception:
pass
try:
_remove_task_result_file(task_id)
except Exception:
pass
for task_file_path in _status_file_paths():
loaded_item = _load_task_payload_from_file(task_file_path)
if loaded_item is None:
try:
task_file_path.unlink()
except Exception:
pass
continue
task_id, payload = loaded_item
payload, _ = _split_task_payload(task_id, payload, persist_legacy_result=True)
if _is_task_expired(payload, now_ts=current_time):
try:
task_file_path.unlink()
except Exception:
pass
try:
_remove_task_result_file(task_id)
except Exception:
pass
_last_cleanup_at = current_time
def recover_tasks_after_restart() -> None:
load_tasks_from_disk()
cleanup_expired_tasks(force=True)
with _tasks_lock:
for task_id, task_payload in list(tasks_db.items()):
if str(task_payload.get("status") or "").strip() not in {"queued", "processing"}:
continue
interrupted_message = "服务已重启,原离线任务已中断,请重新提交。"
task_payload.update(
{
"status": "failed",
"stage": "failed",
"message": interrupted_message,
"error": interrupted_message,
"percentage": 100,
"updated_at": int(time.time()),
}
)
tasks_db[task_id] = task_payload
_write_task_file(task_id, task_payload)
def create_task_record(task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
current_time = int(time.time())
task_payload = dict(payload)
task_payload["task_id"] = task_id
task_payload.setdefault("created_at", current_time)
task_payload["updated_at"] = current_time
status_payload, has_result = _split_task_payload(task_id, task_payload)
with _tasks_lock:
tasks_db[task_id] = status_payload
_write_task_file(task_id, status_payload)
response_payload = dict(status_payload)
if has_result:
response_payload["result"] = _load_task_result_payload(task_id)
return response_payload
def get_task_record(task_id: str, *, include_result: bool = False) -> Optional[dict[str, Any]]:
with _tasks_lock:
task_file_path = _task_file_path(task_id)
if task_file_path.exists():
loaded_item = _load_task_payload_from_file(task_file_path)
if loaded_item is None:
return None
loaded_task_id, task_payload = loaded_item
task_payload, _ = _split_task_payload(
loaded_task_id,
task_payload,
persist_legacy_result=True,
)
if loaded_task_id != task_id or _is_task_expired(task_payload):
return None
tasks_db[task_id] = task_payload
else:
task_payload = tasks_db.get(task_id)
if task_payload is None:
return None
response_payload = dict(task_payload)
if include_result:
result_payload = _load_task_result_payload(task_id)
if result_payload is not None:
response_payload["result"] = result_payload
return response_payload
def update_task_record(task_id: str, payload: dict[str, Any]) -> dict[str, Any]:
current_time = int(time.time())
with _tasks_lock:
task_payload = dict(tasks_db.get(task_id) or {})
task_payload.update(payload)
task_payload["task_id"] = task_id
task_payload.setdefault("created_at", current_time)
task_payload["updated_at"] = current_time
status_payload, has_result = _split_task_payload(task_id, task_payload)
tasks_db[task_id] = status_payload
_write_task_file(task_id, status_payload)
response_payload = dict(status_payload)
if has_result:
response_payload["result"] = _load_task_result_payload(task_id)
return response_payload
load_tasks_from_disk()

View File

@ -0,0 +1,199 @@
# -*- coding: utf-8 -*-
"""Conservative ASR text deduplication and filler cleanup."""
from __future__ import annotations
import re
from app.core.text_cleanup_lexicon import DISCOURSE_FILLER_TOKENS
from app.core.text_cleanup_lexicon import FILLER_TOKENS
from app.core.text_cleanup_lexicon import NUMERIC_STUTTER_CHARS
from app.core.text_cleanup_lexicon import REPEAT_COLLAPSIBLE_DISCOURSE_TOKENS
from app.core.text_cleanup_lexicon import SAFE_DOUBLE_WORDS
from app.core.text_cleanup_lexicon import TERMINAL_FILLER_TOKENS
_PREFIX_STUTTER_CHARS = "这那离超和跟在对把将又还上先后了的"
_PREFIX_STUTTER_PATTERN = re.compile(rf"([{re.escape(_PREFIX_STUTTER_CHARS)}])\1(?!\1)([\u4e00-\u9fffA-Za-z]{{1,3}})")
_REPEATED_PHRASE_PATTERN = re.compile(r"([\u4e00-\u9fffA-Za-z]{2,4})\1")
_WORD_STUTTER_PREFIX_PATTERN = re.compile(r"([\u4e00-\u9fff])\1{1,5}(?=\1[\u4e00-\u9fff])")
_LONG_CHAR_REPEAT_PATTERN = re.compile(r"([\u4e00-\u9fff])\1{2,}")
_INTERNAL_DOUBLE_CHAR_REPEAT_PATTERN = re.compile(r"(?<=[\u4e00-\u9fff])([\u4e00-\u9fff])\1(?=[\u4e00-\u9fff])")
_SHORT_LEADING_STUTTER_TOKEN_PATTERN = re.compile(r"(^|[,。!?;:、,\s])([\u4e00-\u9fff])\2([\u4e00-\u9fff])(?=($|[,。!?;:、,\s]))")
_SHORT_CONFIRMATION_STUTTER_PATTERN = re.compile(r"(^|[,。!?;:、,\s])([\u4e00-\u9fff])\2([\u4e00-\u9fff])(?=(是吧|对吧|对吗|对不对))")
_REPEATED_SEPARATED_PHRASE_PATTERN = re.compile(r"([\u4e00-\u9fffA-Za-z]{1,6})([,、,\s]+)\1(?:\2\1)*")
_REPEATED_SENTENCE_CLAUSE_PATTERN = re.compile(r"([\u4e00-\u9fffA-Za-z0-9]{2,12})([。!?;:]+)(?:\s*\1\2)+")
_REPEATED_SHORT_SENTENCE_CLAUSE_PATTERN = re.compile(
r"([\u4e00-\u9fff])([。!?;:]+)(?:\s*\1\2){2,}(?:\s*\1)?(?=$|[,。!?;:、,\s])"
)
_BOUNDARY_OVERLAP_MIN_CHARS = 2
_BOUNDARY_OVERLAP_MAX_CHARS = 12
_STANDALONE_FILLER_PATTERN = re.compile(rf"(^|[,。!?;:、,\s])({'|'.join(map(re.escape, FILLER_TOKENS))})(?=($|[,。!?;:、,\s]))")
_FILLER_ONLY_PATTERN = re.compile(rf"^[\s,。!?;:、,]*(?:{'|'.join(map(re.escape, FILLER_TOKENS))}[\s,。!?;:、,]*)+$")
_LEADING_FILLER_PREFIX_PATTERN = re.compile(rf"^(?:{'|'.join(map(re.escape, FILLER_TOKENS))})[,、,\s]*")
_INLINE_FILLER_PATTERN = re.compile(r"(?<=[\u4e00-\u9fffA-Za-z0-9])(嗯|呃|啊|(?<!金)额(?!度))(?=[\u4e00-\u9fffA-Za-z0-9])")
_REPEATED_DISCOURSE_FILLER_PATTERN = re.compile(
rf"({'|'.join(map(re.escape, REPEAT_COLLAPSIBLE_DISCOURSE_TOKENS))})(?:[,、,\s]*\1)+"
)
_STANDALONE_DISCOURSE_FILLER_PATTERN = re.compile(
rf"(^|[,。!?;:、,\s])({'|'.join(map(re.escape, DISCOURSE_FILLER_TOKENS))})(?=($|[,。!?;:、,\s]))"
)
_BRIDGE_FILLER_PATTERN = re.compile(r"(和|跟)(?:这个|那个)(和|跟)")
_REPEATED_TOPIC_WITH_DEICTIC_PATTERN = re.compile(r"([\u4e00-\u9fffA-Za-z]{2,6})(这个|那个)\1(?=[这那])")
_TERMINAL_FILLER_PATTERN = re.compile(
rf"(?<=[\u4e00-\u9fffA-Za-z0-9%])(?:{'|'.join(map(re.escape, TERMINAL_FILLER_TOKENS))})(?=($|[,。!?;:、,\s]))"
)
_POST_PUNCT_TERMINAL_FILLER_PATTERN = re.compile(
rf"(?<=[。!?;:,、,])(?:{'|'.join(map(re.escape, TERMINAL_FILLER_TOKENS))})(?=($|[,。!?;:、,\s]))"
)
_CURRENCY_STUTTER_PATTERN = re.compile(r"[¥¥]\s*\d{1,2}\s*[¥¥]\s*(\d{3,})(?=元?|[的档手]|$)")
_CURRENCY_SYMBOL_PATTERN = re.compile(r"[¥¥]\s*(\d+(?:\.\d+)?)")
_HOUSEHOLD_COLLECTION_GLUE_PATTERN = re.compile(r"(收集到)(\d{5,6})(户(?=的(?:一个)?清单))")
_NUMERIC_TOKEN_CHARS = r"0-9零〇○O一幺二两三四五六七八九十百千万亿点"
_REPEATED_NUMERIC_CLAUSE_PATTERN = re.compile(
rf"(?<![A-Za-z0-9])([{_NUMERIC_TOKEN_CHARS}]{{1,8}})([,。!?;:、,\s]+)(?:\1\2){{2,}}"
)
def _is_numeric_like(text: str) -> bool:
return bool(text) and all(character.isdigit() or character in "零〇○O一幺二两三四五六七八九十百千万亿点" for character in text)
def _has_meaningful_overlap(text: str) -> bool:
return bool(text and any(not character.isspace() and character not in ",。!?;:、,.!?;:" for character in text))
def _collapse_internal_double_char(match: re.Match[str]) -> str:
start_index = match.start()
source_text = match.string
pair_text = source_text[start_index:start_index + 2]
if match.group(1) in NUMERIC_STUTTER_CHARS:
return pair_text
if pair_text in SAFE_DOUBLE_WORDS:
return pair_text
return match.group(1)
def _remove_filler_words(text: str) -> str:
normalized_text = str(text or "").strip()
if not normalized_text:
return normalized_text
if _FILLER_ONLY_PATTERN.fullmatch(normalized_text):
return ""
compact_text = _LEADING_FILLER_PREFIX_PATTERN.sub("", normalized_text)
compact_text = _INLINE_FILLER_PATTERN.sub("", compact_text)
compact_text = _BRIDGE_FILLER_PATTERN.sub(lambda match: match.group(1) if match.group(1) == match.group(2) else match.group(2), compact_text)
compact_text = _REPEATED_TOPIC_WITH_DEICTIC_PATTERN.sub(lambda match: match.group(1), compact_text)
compact_text = _REPEATED_DISCOURSE_FILLER_PATTERN.sub(lambda match: match.group(1), compact_text)
compact_text = _STANDALONE_FILLER_PATTERN.sub(lambda match: match.group(1), compact_text)
compact_text = _STANDALONE_DISCOURSE_FILLER_PATTERN.sub(lambda match: match.group(1), compact_text)
compact_text = _TERMINAL_FILLER_PATTERN.sub("", compact_text)
compact_text = _POST_PUNCT_TERMINAL_FILLER_PATTERN.sub("", compact_text)
compact_text = re.sub(r"([。!?;:])[。!?;:]+", r"\1", compact_text)
compact_text = re.sub(r"([。!?;:])[,、,]+", r"\1", compact_text)
compact_text = re.sub(r"[,、,\s]{2,}", ",", compact_text)
compact_text = re.sub(r"^[,、,\s]+", "", compact_text)
compact_text = re.sub(r"[,、,\s]+([。!?;:])", r"\1", compact_text)
compact_text = re.sub(r"[,、,\s]+$", "", compact_text)
if not compact_text or re.fullmatch(r"[\s,。!?;:、,]*", compact_text):
return ""
return compact_text
def _repair_household_collection_count(match: re.Match[str]) -> str:
prefix, digits, suffix = match.groups()
following_context = match.string[match.end():match.end() + 100]
# Meeting reports often say "collected N households, estimated M households,
# close to 50%". If the glued count has a plausible suffix matching that
# ratio, keep the suffix and drop the noisy leading recognition artifact.
if "50" in following_context:
reference_counts = [
int(value)
for value in re.findall(r"(?:大概有|约|预计|测算[^,。!?;:]{0,10}?有)(\d{1,4})户", following_context)
]
for suffix_len in range(min(4, len(digits) - 1), 1, -1):
candidate = int(digits[-suffix_len:])
if candidate <= 0:
continue
if any(0.35 <= reference_count / candidate <= 0.65 for reference_count in reference_counts):
return f"{prefix}{candidate}{suffix}"
return match.group(0)
def _repair_numeric_asr_artifacts(text: str) -> str:
repaired_text = _CURRENCY_STUTTER_PATTERN.sub(lambda match: f"{match.group(1)}元", text)
repaired_text = _CURRENCY_SYMBOL_PATTERN.sub(lambda match: f"{match.group(1)}元", repaired_text)
repaired_text = _HOUSEHOLD_COLLECTION_GLUE_PATTERN.sub(_repair_household_collection_count, repaired_text)
return repaired_text
def _collapse_repeated_numeric_clauses(text: str) -> str:
return _REPEATED_NUMERIC_CLAUSE_PATTERN.sub(lambda match: f"{match.group(1)}{match.group(2)}", text)
def _collapse_repeated_short_sentence_clauses(text: str) -> str:
return _REPEATED_SHORT_SENTENCE_CLAUSE_PATTERN.sub(lambda match: f"{match.group(1)}{match.group(2)}", text)
def deduplicate_asr_text(text: str) -> str:
normalized_text = str(text or "").strip()
if not normalized_text:
return normalized_text
previous_text = None
while previous_text != normalized_text:
previous_text = normalized_text
normalized_text = _PREFIX_STUTTER_PATTERN.sub(lambda match: f"{match.group(1)}{match.group(2)}", normalized_text)
def _collapse_repeated_phrase(match: re.Match[str]) -> str:
phrase = match.group(1)
if _is_numeric_like(phrase):
return match.group(0)
return phrase
normalized_text = _REPEATED_PHRASE_PATTERN.sub(_collapse_repeated_phrase, normalized_text)
normalized_text = _WORD_STUTTER_PREFIX_PATTERN.sub(lambda match: match.group(0) if match.group(1) in NUMERIC_STUTTER_CHARS else "", normalized_text)
normalized_text = _LONG_CHAR_REPEAT_PATTERN.sub(lambda match: match.group(0) if match.group(1) in NUMERIC_STUTTER_CHARS else match.group(1), normalized_text)
normalized_text = _INTERNAL_DOUBLE_CHAR_REPEAT_PATTERN.sub(_collapse_internal_double_char, normalized_text)
normalized_text = _SHORT_LEADING_STUTTER_TOKEN_PATTERN.sub(lambda match: f"{match.group(1)}{match.group(2)}{match.group(3)}", normalized_text)
def _collapse_confirmation(match: re.Match[str]) -> str:
repeated_pair = f"{match.group(2)}{match.group(2)}"
if match.group(2) in NUMERIC_STUTTER_CHARS or repeated_pair in SAFE_DOUBLE_WORDS:
return match.group(0)
return f"{match.group(1)}{match.group(2)}{match.group(3)}"
normalized_text = _SHORT_CONFIRMATION_STUTTER_PATTERN.sub(_collapse_confirmation, normalized_text)
def _collapse_separated(match: re.Match[str]) -> str:
phrase = match.group(1)
if phrase.isascii() and len(phrase) == 1:
return match.group(0)
if _is_numeric_like(phrase):
return match.group(0)
return phrase
normalized_text = _REPEATED_SEPARATED_PHRASE_PATTERN.sub(_collapse_separated, normalized_text)
normalized_text = _REPEATED_SENTENCE_CLAUSE_PATTERN.sub(lambda match: match.group(0) if _is_numeric_like(match.group(1)) else f"{match.group(1)}{match.group(2)}", normalized_text)
normalized_text = _collapse_repeated_short_sentence_clauses(normalized_text)
normalized_text = _collapse_repeated_numeric_clauses(normalized_text)
normalized_text = _remove_filler_words(normalized_text)
normalized_text = _repair_numeric_asr_artifacts(normalized_text)
return normalized_text
def trim_segment_boundary_overlap(previous_text: str, current_text: str) -> tuple[str, bool]:
normalized_previous = str(previous_text or "").strip()
normalized_current = str(current_text or "").strip()
if not normalized_previous or not normalized_current:
return normalized_current, False
max_overlap_chars = min(len(normalized_previous), len(normalized_current), _BOUNDARY_OVERLAP_MAX_CHARS)
for overlap_chars in range(max_overlap_chars, _BOUNDARY_OVERLAP_MIN_CHARS - 1, -1):
overlap_suffix = normalized_previous[-overlap_chars:]
overlap_prefix = normalized_current[:overlap_chars]
if overlap_suffix != overlap_prefix:
continue
if not _has_meaningful_overlap(overlap_prefix):
continue
return normalized_current[overlap_chars:].lstrip(), True
return normalized_current, False

View File

@ -0,0 +1,64 @@
# -*- coding: utf-8 -*-
"""Conservative ASR text cleanup lexicons."""
from __future__ import annotations
SAFE_DOUBLE_WORDS = {
"哥哥",
"叔叔",
"爸爸",
"妈妈",
"奶奶",
"爷爷",
"姐姐",
"弟弟",
"妹妹",
"伯伯",
"姑姑",
"舅舅",
"星星",
"猩猩",
}
NUMERIC_STUTTER_CHARS = set("零〇○O一幺二两三四五六七八九十百千万亿点")
FILLER_TOKENS = (
"啊",
"呢",
"吧",
"哦",
"嗯",
"呃",
"哈",
"哇",
"呀",
"哎",
"诶",
"欸",
"额",
)
DISCOURSE_FILLER_TOKENS = (
"那个",
"这个",
"然后",
"所以说",
"其实",
"那么",
)
REPEAT_COLLAPSIBLE_DISCOURSE_TOKENS = (
"那个",
"这个",
"然后",
"就是",
"所以",
"所以说",
"其实",
"那么",
)
TERMINAL_FILLER_TOKENS = (
"哈",
"嗯",
)

View File

@ -0,0 +1,8 @@
# -*- coding: utf-8 -*-
"""
基础设施层 - 提供底层通用功能
"""
from .model_utils import resolve_model_path
__all__ = ["resolve_model_path"]

View File

@ -0,0 +1,48 @@
# -*- coding: utf-8 -*-
"""
模型工具模块 - 提供模型路径解析等通用功能
"""
import logging
from pathlib import Path
from typing import Optional
from app.core.config import settings
logger = logging.getLogger(__name__)
def resolve_model_path(model_id: Optional[str]) -> str:
"""将模型 ID 解析为本地模型路径(如果存在)
本项目默认模型目录结构:
./models/{publisher}/{model_name}/
如果本地模型存在,返回本地路径;否则返回原始 model_id
"""
if not model_id:
raise ValueError("model_id 不能为空")
# 项目内扁平化模型根目录
local_path = Path(settings.MODELSCOPE_PATH) / model_id
if local_path.exists() and local_path.is_dir():
resolved = str(local_path)
logger.info(f"模型 {model_id} 使用本地缓存: {resolved}")
return resolved
# 兼容历史错误配置:MODELSCOPE_CACHE 指向了 models 目录时,
# ModelScope 会生成 /models/models/{publisher}/{model_name}。
legacy_nested_path = Path(settings.MODELSCOPE_PATH) / "models" / model_id
if legacy_nested_path.exists() and legacy_nested_path.is_dir():
resolved = str(legacy_nested_path)
logger.warning(
"模型 %s 命中历史嵌套缓存: %s。建议迁移到 %s",
model_id,
resolved,
local_path,
)
return resolved
logger.warning(f"模型 {model_id} 本地缓存不存在,将在运行时下载")
return model_id

283
app/main.py 100644
View File

@ -0,0 +1,283 @@
# -*- coding: utf-8 -*-
"""
FastAPI应用创建和配置
"""
import warnings
import asyncio
import os
import logging
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException
from fastapi.exceptions import RequestValidationError
from fastapi_offline import FastAPIOffline
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from .core.config import settings
from .core.exceptions import (
APIException,
api_exception_handler,
general_exception_handler,
http_exception_handler,
validation_exception_handler,
)
from .core.logging import setup_logging, get_worker_id
from .core.executor import shutdown_executor
from .api.v1 import api_router
from .utils.boot_events import emit_boot_event
# 忽略 Pydantic V2 兼容性警告
warnings.filterwarnings("ignore", message="Valid config keys have changed in V2")
warnings.filterwarnings("ignore", message=".*has conflict with protected namespace.*")
warnings.filterwarnings("ignore", category=UserWarning, module="pydantic")
logger = logging.getLogger(__name__)
async def cleanup_task_state_loop():
"""定期清理过期离线任务状态文件。"""
while True:
await asyncio.sleep(600)
try:
from .core.task_store import cleanup_expired_tasks
cleanup_expired_tasks(force=True)
except asyncio.CancelledError:
raise
except Exception as exc:
logger.warning("定期清理离线任务状态失败: %s", exc)
def cleanup_temp_directory():
"""清理临时目录中的旧文件"""
import time
temp_dir = settings.TEMP_DIR
if not os.path.exists(temp_dir):
return
# 清理超过 1 小时的临时文件
max_age_seconds = 3600
current_time = time.time()
cleaned_count = 0
try:
for filename in os.listdir(temp_dir):
filepath = os.path.join(temp_dir, filename)
if os.path.isfile(filepath):
file_age = current_time - os.path.getmtime(filepath)
if file_age > max_age_seconds:
try:
os.remove(filepath)
cleaned_count += 1
except Exception:
pass
if cleaned_count > 0:
logger.info(f"已清理 {cleaned_count} 个过期临时文件")
except Exception as e:
logger.warning(f"清理临时目录时出错: {e}")
@asynccontextmanager
async def lifespan(app: FastAPI):
"""应用生命周期管理"""
workers = int(os.getenv("WORKERS", "1"))
worker_id = get_worker_id()
task_cleanup_task = None
# 启动时
logger.info(f"Worker [{worker_id}] 启动中...")
emit_boot_event("phase_start", phase="worker", total=1, message=f"Worker [{worker_id}] 启动中")
# 清理旧的临时文件(仅主 Worker 执行)
if worker_id == 0:
cleanup_temp_directory()
try:
from .core.task_store import recover_tasks_after_restart
recover_tasks_after_restart()
except Exception as exc:
logger.warning("恢复离线任务状态失败: %s", exc)
task_cleanup_task = asyncio.create_task(cleanup_task_state_loop())
if settings.SPEAKER_DB_ENABLED:
try:
from .core.database import pg_speaker_db
await pg_speaker_db.connect()
except Exception as exc:
logger.warning(
"声纹数据库连接失败,将保留原说话人编号并禁用声纹注册接口: %s",
exc,
)
from .utils.model_loader import (
preload_models,
verify_required_models_integrity,
)
integrity_result = verify_required_models_integrity()
if integrity_result["invalid_models"]:
emit_boot_event("error", phase="integrity", message="required model integrity check failed")
raise RuntimeError("required model integrity check failed")
logger.info(f"Worker [{worker_id}] 正在加载模型...")
preload_result = preload_models()
asr_results = preload_result.get("asr_models", {})
loaded_count = sum(1 for r in asr_results.values() if r.get("loaded"))
total_count = len(asr_results)
logger.info(f"Worker [{worker_id}] 模型加载完成: {loaded_count}/{total_count}")
failed_asr_models = {
model_id: status.get("error")
for model_id, status in asr_results.items()
if not status.get("loaded") and status.get("error")
}
try:
from .services.asr.model_plan import get_active_qwen_model, load_supported_model_ids
active_qwen_model = get_active_qwen_model(load_supported_model_ids())
except Exception as exc:
logger.error(f"Worker [{worker_id}] 无法解析当前应启用的 Qwen 模型: {exc}")
emit_boot_event("error", phase="preload", message=f"无法解析当前应启用的 Qwen 模型: {exc}")
raise
active_qwen_status = asr_results.get(active_qwen_model, {})
if not active_qwen_status.get("loaded"):
qwen_error = active_qwen_status.get("error") or "unknown error"
logger.error(
f"Worker [{worker_id}] Qwen 主模型预加载失败,拒绝启动: {active_qwen_model}, error={qwen_error}"
)
emit_boot_event(
"error",
phase="preload",
message=f"Qwen 主模型预加载失败,拒绝启动: {active_qwen_model}, error={qwen_error}",
)
raise RuntimeError(
f"required qwen model preload failed: {active_qwen_model}: {qwen_error}"
)
if failed_asr_models:
if loaded_count == 0:
logger.error(f"Worker [{worker_id}] ASR模型预加载失败详情: {failed_asr_models}")
emit_boot_event("error", phase="preload", message=f"ASR模型预加载失败详情: {failed_asr_models}")
raise RuntimeError(f"ASR model preload failed: {failed_asr_models}")
logger.warning(
f"Worker [{worker_id}] 部分ASR模型预加载失败,将以可用模型继续启动: {failed_asr_models}"
)
emit_boot_event(
"warning",
phase="preload",
message=f"部分ASR模型预加载失败,将以可用模型继续启动: {failed_asr_models}",
)
if settings.SPEAKER_DB_ENABLED:
try:
from .core.database import pg_speaker_db
from .services.speaker_registry import get_speaker_registry_service
if pg_speaker_db.is_connected:
logger.info(f"Worker [{worker_id}] 正在预加载声纹识别模型...")
get_speaker_registry_service().ensure_loaded()
except Exception as exc:
logger.warning("声纹识别模型预加载失败,将在首次请求时重试: %s", exc)
logger.info(f"Worker [{worker_id}] 已就绪")
emit_boot_event("ready", phase="worker", message=f"Worker [{worker_id}] 已就绪")
yield
# 关闭时
if task_cleanup_task is not None:
task_cleanup_task.cancel()
try:
await task_cleanup_task
except asyncio.CancelledError:
pass
if settings.SPEAKER_DB_ENABLED:
try:
from .core.database import pg_speaker_db
await pg_speaker_db.close()
except Exception as exc:
logger.warning("关闭声纹数据库连接时出错: %s", exc)
logger.info(f"Worker [{worker_id}] 正在关闭推理线程池...")
shutdown_executor()
logger.info(f"Worker [{worker_id}] 已关闭")
def create_app() -> FastAPI:
"""创建FastAPI应用"""
# 设置日志
setup_logging()
app = FastAPIOffline(
title=settings.APP_NAME,
description=settings.APP_DESCRIPTION,
version=settings.APP_VERSION,
docs_url=settings.docs_url,
redoc_url=settings.redoc_url,
lifespan=lifespan, # 添加生命周期管理
)
# 添加CORS中间件
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# 注册异常处理器
app.add_exception_handler(APIException, api_exception_handler)
app.add_exception_handler(HTTPException, http_exception_handler)
app.add_exception_handler(RequestValidationError, validation_exception_handler)
app.add_exception_handler(Exception, general_exception_handler)
# 注册静态文件服务(用于临时文件)
app.mount("/tmp", StaticFiles(directory=settings.TEMP_DIR), name="temp_files")
app.mount(
"/test-web",
StaticFiles(directory=str(settings.BASE_DIR / "test_web"), html=True),
name="test_web",
)
# 注册API路由
app.include_router(api_router)
# 根路径
@app.get("/", summary="根路径", description="API服务根路径")
async def root():
return {
"message": settings.APP_NAME,
"version": settings.APP_VERSION,
"description": settings.APP_DESCRIPTION,
"endpoints": {
# 阿里云兼容 API
"asr": "/stream/v1/asr",
"asr_models": "/stream/v1/asr/models",
"asr_health": "/stream/v1/asr/health",
"ws_asr": "/ws/v1/asr/qwen",
"ws_qwen3_asr": "/ws/v1/asr/qwen",
# OpenAI 兼容 API
"openai_models": "/v1/models",
"openai_transcriptions": "/v1/audio/transcriptions",
# 独立会议离线 / 声纹管理 API
"meeting_files": "/api/v1/files",
"meeting_transcriptions": "/api/v1/asr/transcriptions",
"speakers": "/api/v1/speakers",
"test_web": "/test-web/",
# 文档
"docs": settings.docs_url or "禁用",
},
}
return app
# 创建全局应用实例
app = create_app()

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
数据模型模块
包含API的请求和响应模型定义
"""

321
app/models/asr.py 100644
View File

@ -0,0 +1,321 @@
# -*- coding: utf-8 -*-
"""
ASR数据模型
定义语音识别相关的请求和响应模型
"""
from typing import Optional, List, Union
from pydantic import BaseModel, Field
from .common import (
SampleRate,
BaseResponse,
HealthCheckResponse,
ErrorResponse,
)
# ============= 请求模型 =============
class ASRQueryParams(BaseModel):
"""ASR接口查询参数模型"""
model: Optional[str] = Field(
default=None,
description="可选。离线 ASR 模型 ID;不传则使用服务当前默认模型,如 qwen3-asr-0.6b 或 qwen3-asr-1.7b",
max_length=128,
)
audio_address: Optional[str] = Field(
default=None,
description="音频/视频文件地址,支持 HTTP/HTTPS URL、file:// 或服务端本地路径,格式自动识别",
max_length=512,
)
sample_rate: Optional[SampleRate] = Field(
default=SampleRate.RATE_16000,
description=f"音频采样率(Hz)。支持: {', '.join(map(str, SampleRate.get_enums()))}",
)
enable_speaker_diarization: Optional[bool] = Field(
default=True,
description="是否启用说话人分离。启用后响应会包含 speaker_id",
)
enable_speaker_identification: Optional[bool] = Field(
default=True,
description="是否匹配已注册声纹库。仅在 enable_speaker_diarization=true 时生效",
)
enable_text_cleanup: Optional[bool] = Field(
default=True,
description="是否启用文本去重和口头语清理",
)
word_timestamps: Optional[bool] = Field(
default=False,
description="是否返回字词级时间戳(默认关闭;Qwen CUDA vLLM / CPU Rust 会在启用时自动调用 forced aligner)",
)
vocabulary_id: Optional[str] = Field(
default=None,
description="热词字符串,格式:热词1 权重1 热词2 权重2(如:阿里巴巴 20 腾讯 15)",
max_length=512,
)
# ============= 响应模型 =============
class WordToken(BaseModel):
"""字词级时间戳信息"""
text: str = Field(
...,
description="字词文本",
)
start_time: float = Field(
...,
description="开始时间(秒)",
)
end_time: float = Field(
...,
description="结束时间(秒)",
)
model_config = {
"json_schema_extra": {
"example": {
"text": "今",
"start_time": 0.0,
"end_time": 0.15,
}
}
}
class ASRSegment(BaseModel):
"""ASR 识别分段结果"""
text: str = Field(
...,
description="该段识别文本",
)
start_time: float = Field(
...,
description="段落开始时间(秒)",
)
end_time: float = Field(
...,
description="段落结束时间(秒)",
)
speaker_id: Optional[str] = Field(
default=None,
description="说话人ID(如 说话人1),仅启用说话人分离时返回",
)
word_tokens: Optional[List[WordToken]] = Field(
default=None,
description="字词级时间戳(仅启用 word_timestamps 且模型支持时返回)",
)
model_config = {
"json_schema_extra": {
"example": {
"text": "今天天气不错。",
"start_time": 0.0,
"end_time": 2.5,
"speaker_id": "说话人1",
"word_tokens": [
{"text": "今", "start_time": 0.0, "end_time": 0.15},
{"text": "天", "start_time": 0.15, "end_time": 0.35},
],
}
}
}
class ASRSuccessResponse(BaseResponse):
"""ASR成功响应模型"""
result: str = Field(
...,
description="识别结果文本(完整)",
max_length=100000,
)
segments: Optional[List[ASRSegment]] = Field(
default=None,
description="分段识别结果(含时间戳),仅长音频分段识别时返回",
)
duration: Optional[float] = Field(
default=None,
description="音频总时长(秒)",
)
model_config = {
"json_schema_extra": {
"example": {
"task_id": "cf7b0c5339244ee29cd4e43fb97f1234",
"result": "今天天气不错。明天可能会下雨。",
"segments": [
{"text": "今天天气不错。", "start_time": 0.0, "end_time": 2.5, "speaker_id": "说话人1"},
{"text": "明天可能会下雨。", "start_time": 3.2, "end_time": 5.8, "speaker_id": "说话人2"},
],
"duration": 5.8,
"status": 200,
"message": "SUCCESS",
}
}
}
class ASRErrorResponse(ErrorResponse):
"""ASR错误响应模型"""
result: str = Field(default="", description="识别结果(错误时为空)")
model_config = {
"json_schema_extra": {
"example": {
"task_id": "8bae3613dfc54ebfa811a17d8a7a1234",
"result": "",
"status": 40000001,
"message": "Gateway:ACCESS_DENIED:The token 'invalid_token' is invalid!",
}
}
}
class ASRHealthCheckResponse(HealthCheckResponse):
"""ASR健康检查响应模型"""
model_config = {
"protected_namespaces": (),
"json_schema_extra": {
"example": {
"status": "healthy",
"model_loaded": True,
"device": "cuda:0",
"version": "1.0.0",
"message": "ASR service is running normally",
"loaded_models": ["qwen3-asr-1.7b"],
"memory_usage": {
"gpu_memory_used": "2.1GB",
"gpu_memory_total": "8.0GB",
},
"accelerator": {
"vendor": "nvidia",
"runtime": "cuda",
"device": "cuda:0",
"device_count": 1,
},
},
},
}
model_loaded: bool = Field(..., description="模型是否已加载")
device: str = Field(..., description="推理设备")
loaded_models: Optional[List[str]] = Field(default=[], description="已加载的模型列表")
memory_usage: Optional[dict] = Field(default=None, description="内存使用情况")
accelerator: Optional[dict] = Field(default=None, description="加速器信息")
# ============= 模型相关 =============
class ASRDeclaredEntryInfo(BaseModel):
"""声明式 ASR 条目信息,可表示离线模型或 realtime capability。"""
id: str = Field(..., description="模型id")
kind: str = Field(..., description="条目类型:model 或 capability")
name: str = Field(..., description="模型名称")
engine: str = Field(..., description="引擎类型")
description: str = Field(..., description="模型描述")
languages: List[str] = Field(..., description="支持的语言列表")
default: bool = Field(default=False, description="是否为默认模型")
supports_realtime: bool = Field(default=False, description="是否支持实时识别")
offline_model: Optional[dict] = Field(default=None, description="离线模型信息")
realtime_model: Optional[dict] = Field(default=None, description="实时模型信息")
model_config = {
"json_schema_extra": {
"example": {
"id": "qwen3-asr-1.7b",
"kind": "model",
"name": "Qwen3-ASR-1.7B",
"engine": "qwen3",
"description": "多语言离线语音识别模型",
"languages": ["zh", "en"],
"default": True,
"supports_realtime": True,
"offline_model": {
"path": "Qwen/Qwen3-ASR-1.7B",
"exists": True,
},
"realtime_model": None,
}
}
}
class ASRRuntimeInfo(BaseModel):
"""运行时视角的模型加载状态。"""
loaded_model_ids: List[str] = Field(default_factory=list, description="当前已加载模型 ID 列表")
loaded_count: int = Field(..., description="已加载模型数量")
default_offline_model_id: Optional[str] = Field(default=None, description="当前默认离线模型 ID")
model_config = {
"json_schema_extra": {
"example": {
"loaded_model_ids": ["qwen3-asr-1.7b"],
"loaded_count": 1,
"default_offline_model_id": "qwen3-asr-1.7b",
}
}
}
class ASRModelsResponse(BaseModel):
"""ASR 模型列表响应,分离声明视角与运行时视角。"""
declared_entries: List[ASRDeclaredEntryInfo] = Field(..., description="声明的模型与 capability 列表")
declared_count: int = Field(..., description="声明条目总数")
runtime: ASRRuntimeInfo = Field(..., description="运行时加载状态")
model_config = {
"json_schema_extra": {
"example": {
"declared_entries": [
{
"id": "qwen3-asr-1.7b",
"kind": "model",
"name": "Qwen3-ASR-1.7B",
"engine": "qwen3",
"description": "多语言离线语音识别模型",
"languages": ["zh", "en"],
"default": True,
"supports_realtime": True,
"offline_model": {
"path": "Qwen/Qwen3-ASR-1.7B",
"exists": True,
},
"realtime_model": None,
}
],
"declared_count": 2,
"runtime": {
"loaded_model_ids": ["qwen3-asr-1.7b"],
"loaded_count": 1,
"default_offline_model_id": "qwen3-asr-1.7b",
},
}
}
}
# ============= 联合响应类型 =============
ASRResponse = Union[ASRSuccessResponse, ASRErrorResponse]

View File

@ -0,0 +1,64 @@
# -*- coding: utf-8 -*-
"""
通用数据模型
定义通用的枚举、基础模型等
"""
from pydantic import BaseModel, Field
from enum import Enum
class AudioFormat(str, Enum):
"""支持的音频格式"""
PCM = "pcm"
WAV = "wav"
OPUS = "opus"
SPEEX = "speex"
AMR = "amr"
MP3 = "mp3"
AAC = "aac"
M4A = "m4a"
FLAC = "flac"
OGG = "ogg"
@classmethod
def get_enums(cls):
return [e.value for e in cls]
class SampleRate(int, Enum):
"""支持的采样率"""
RATE_8000 = 8000
RATE_16000 = 16000
RATE_22050 = 22050
RATE_24000 = 24000
@classmethod
def get_enums(cls):
return [e.value for e in cls]
class BaseResponse(BaseModel):
"""基础响应模型"""
task_id: str = Field(..., description="任务ID")
status: int = Field(..., description="状态码")
message: str = Field(..., description="响应消息")
class HealthCheckResponse(BaseModel):
"""健康检查响应模型"""
status: str = Field(description="服务状态", examples=["healthy"])
version: str = Field(description="服务版本", examples=["1.0.0"])
message: str = Field(
description="状态消息", examples=["Service is running normally"]
)
class ErrorResponse(BaseResponse):
"""错误响应模型"""
result: str = Field("", description="结果内容")

View File

@ -0,0 +1,135 @@
# -*- coding: utf-8 -*-
"""
WebSocket ASR 数据模型 - 阿里云协议
"""
from typing import Optional, Dict, Any, Union, List
from pydantic import BaseModel, Field, field_validator
import uuid
class AliyunASRWSHeader(BaseModel):
"""阿里云WebSocket ASR消息头部"""
message_id: str = Field(..., description="消息ID,32位唯一ID")
task_id: str = Field(..., description="任务ID,32位唯一ID")
namespace: str = Field(..., description="命名空间,固定为SpeechTranscriber")
name: str = Field(..., description="消息名称")
appkey: Optional[str] = Field(None, description="应用密钥")
status: Optional[int] = Field(None, description="状态码")
status_text: Optional[str] = Field(None, description="状态文本")
status_message: Optional[str] = Field(None, description="状态消息")
@staticmethod
def generate_message_id() -> str:
"""生成32位消息ID"""
return str(uuid.uuid4()).replace("-", "")[:32]
class AliyunStartTranscriptionPayload(BaseModel):
"""StartTranscription 消息负载"""
format: str = Field(default="pcm", description="音频格式: pcm, wav, opus, speex, amr, mp3, aac")
sample_rate: int = Field(default=16000, description="音频采样率: 8000/16000")
enable_intermediate_result: bool = Field(default=True, description="是否返回中间识别结果")
enable_punctuation_prediction: bool = Field(default=True, description="是否在后处理中添加标点")
enable_inverse_text_normalization: bool = Field(default=True, description="是否将中文数字转为阿拉伯数字")
customization_id: Optional[str] = Field(None, description="自学习模型ID")
vocabulary_id: Optional[str] = Field(None, description="定制泛热词ID")
max_sentence_silence: int = Field(default=800, ge=200, le=2000, description="语音断句检测阈值(ms)")
enable_words: bool = Field(default=False, description="是否开启返回词信息")
disfluency: bool = Field(default=False, description="过滤语气词")
speech_noise_threshold: Optional[float] = Field(None, ge=-1.0, le=1.0, description="噪音参数阈值")
enable_semantic_sentence_detection: bool = Field(default=False, description="是否开启语义断句")
@field_validator("format")
@classmethod
def validate_format(cls, v):
supported_formats = ["pcm", "wav", "opus", "speex", "amr", "mp3", "aac"]
if v.lower() not in supported_formats:
raise ValueError(f"不支持的音频格式: {v}")
return v.lower()
@field_validator("sample_rate")
@classmethod
def validate_sample_rate(cls, v):
supported_rates = [8000, 16000]
if v not in supported_rates:
raise ValueError(f"不支持的采样率: {v}")
return v
class AliyunWordInfo(BaseModel):
"""词信息"""
text: str = Field("", description="文本")
startTime: int = Field(0, description="词开始时间(ms)")
endTime: int = Field(0, description="词结束时间(ms)")
class AliyunTranscriptionResultPayload(BaseModel):
"""识别结果负载"""
session_id: Optional[str] = Field(None, description="会话ID")
index: Optional[int] = Field(None, description="句子编号,从1开始递增")
time: Optional[int] = Field(None, description="已处理的音频时长(ms)")
begin_time: Optional[int] = Field(None, description="句子开始时间(ms)")
result: Optional[str] = Field(None, description="识别结果文本")
confidence: Optional[float] = Field(None, description="置信度[0.0,1.0]")
words: Optional[List[AliyunWordInfo]] = Field(None, description="词信息列表")
status: Optional[int] = Field(None, description="状态码")
class AliyunStashResult(BaseModel):
"""暂存结果(语义断句)"""
sentenceId: int = Field(0, description="句子编号")
beginTime: int = Field(0, description="句子开始时间(ms)")
text: str = Field("", description="转写内容")
currentTime: int = Field(0, description="当前处理时间(ms)")
class AliyunASRWSMessage(BaseModel):
"""阿里云WebSocket ASR消息"""
header: AliyunASRWSHeader = Field(..., description="消息头部")
payload: Optional[
Union[
AliyunStartTranscriptionPayload,
AliyunTranscriptionResultPayload,
Dict[str, Any],
]
] = Field(None, description="消息负载")
class AliyunASRNamespace:
"""阿里云ASR命名空间"""
SPEECH_TRANSCRIBER = "SpeechTranscriber"
class AliyunASRMessageName:
"""阿里云ASR消息名称"""
START_TRANSCRIPTION = "StartTranscription"
STOP_TRANSCRIPTION = "StopTranscription"
TRANSCRIPTION_STARTED = "TranscriptionStarted"
SENTENCE_BEGIN = "SentenceBegin"
TRANSCRIPTION_RESULT_CHANGED = "TranscriptionResultChanged"
SENTENCE_END = "SentenceEnd"
TRANSCRIPTION_COMPLETED = "TranscriptionCompleted"
TASK_FAILED = "TaskFailed"
class AliyunASRStatus:
"""阿里云ASR状态码"""
SUCCESS = 20000000
TASK_FAILED = 40000000
INVALID_PARAMETER = 40000001
MESSAGE_INVALID = 40000002
AUTHENTICATION_FAILED = 40100005
QUOTA_EXCEEDED = 40300016
INTERNAL_ERROR = 50000000
SERVICE_UNAVAILABLE = 50300018

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
服务层模块
包含业务逻辑和模型管理
"""

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
ASR服务模块
包含语音识别相关的服务和引擎
"""

View File

@ -0,0 +1,24 @@
# -*- coding: utf-8 -*-
"""Audio-related request validation helpers."""
from __future__ import annotations
from typing import Optional
from ...core.exceptions import InvalidParameterException
from ...models.common import SampleRate
SUPPORTED_SAMPLE_RATES = SampleRate.get_enums()
def validate_sample_rate(rate: Optional[int]) -> int:
"""Validate input sample rate."""
if not rate:
return 16000
if rate not in SUPPORTED_SAMPLE_RATES:
raise InvalidParameterException(
f"不支持的采样率: {rate}。支持的采样率: {', '.join(map(str, SUPPORTED_SAMPLE_RATES))}"
)
return rate

View File

@ -0,0 +1,41 @@
# -*- coding: utf-8 -*-
"""
ASR引擎模块
支持多种ASR引擎实现
"""
# 基础类和数据类
from .base import (
BaseASREngine,
RealTimeASREngine,
WordToken,
ASRSegmentResult,
ASRFullResult,
ASRRawResult,
)
# 全局模型管理
from .global_models import (
get_global_vad_model,
get_global_punc_model,
get_global_punc_realtime_model,
get_punc_inference_lock,
get_punc_realtime_inference_lock,
)
__all__ = [
# 基础类
"BaseASREngine",
"RealTimeASREngine",
# 数据类
"WordToken",
"ASRSegmentResult",
"ASRFullResult",
"ASRRawResult",
# 全局模型管理
"get_global_vad_model",
"get_global_punc_model",
"get_global_punc_realtime_model",
"get_punc_inference_lock",
"get_punc_realtime_inference_lock",
]

View File

@ -0,0 +1,785 @@
# -*- coding: utf-8 -*-
"""
ASR引擎基础模块
包含抽象基类和数据类定义
"""
import time
import logging
from typing import Optional, Dict, List, Any, Callable
from abc import ABC, abstractmethod
from dataclasses import dataclass
from app.core.config import settings
from app.core.exceptions import DefaultServerErrorException
from app.core.text_cleanup import deduplicate_asr_text
from app.core.text_cleanup import trim_segment_boundary_overlap
from app.core.hotword_resolver import apply_hotword_rules
from app.core.logging import log_inference_metrics
from app.utils.audio import get_audio_duration
logger = logging.getLogger(__name__)
def _log_asr_stage_timing(
stage: str,
duration_ms: float,
*,
task_id: Optional[str] = None,
model_id: Optional[str] = None,
audio_duration_sec: Optional[float] = None,
**extra: Any,
) -> None:
"""Record structured timing for one ASR processing stage."""
payload: Dict[str, Any] = {
"event": "asr_stage_timing",
"stage": stage,
"duration_ms": round(duration_ms, 2),
}
if task_id:
payload["task_id"] = task_id
if model_id:
payload["model_id"] = model_id
if audio_duration_sec is not None:
payload["audio_duration_sec"] = round(audio_duration_sec, 2)
if audio_duration_sec > 0:
payload["rtf"] = round((duration_ms / 1000) / audio_duration_sec, 4)
payload.update(extra)
logger.info("ASR阶段耗时", extra=payload)
@dataclass
class WordToken:
"""字词级时间戳信息"""
text: str # 字词文本
start_time: float # 开始时间(秒)
end_time: float # 结束时间(秒)
@dataclass
class ASRSegmentResult:
"""ASR 分段识别结果"""
text: str # 该段识别文本
start_time: float # 开始时间(秒)
end_time: float # 结束时间(秒)
speaker_id: Optional[str] = None # 说话人ID(多说话人模式)
speaker_name: Optional[str] = None # 已注册声纹命中后的显示名称
user_id: Optional[str] = None # 已注册声纹命中后的业务用户编号
speaker_embedding: Optional[Any] = None # 片段声纹向量,仅供服务端匹配使用
word_tokens: Optional[List[WordToken]] = None # 字词级时间戳(可选)
@dataclass
class ASRFullResult:
"""ASR 完整识别结果(支持长音频)"""
text: str # 完整识别文本
segments: List[ASRSegmentResult] # 分段结果
duration: float # 音频总时长(秒)
@dataclass
class ASRRawResult:
"""ASR 原始识别结果(包含时间戳)"""
text: str # 完整识别文本
segments: List[ASRSegmentResult] # 分段结果(从 VAD 时间戳解析)
class BaseASREngine(ABC):
"""基础ASR引擎抽象基类"""
@abstractmethod
def transcribe_file(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
enable_vad: bool = False,
sample_rate: int = 16000,
) -> str:
"""转录音频文件"""
pass
@abstractmethod
def transcribe_file_with_vad(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = True,
enable_itn: bool = True,
sample_rate: int = 16000,
**kwargs,
) -> ASRRawResult:
"""使用 VAD 转录音频文件,返回带时间戳分段的结果
Args:
audio_path: 音频文件路径
hotwords: 热词/上下文提示
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
**kwargs: 额外参数(如 word_timestamps 字词级时间戳)
Returns:
ASRRawResult 包含文本和分段信息
"""
pass
def transcribe_long_audio(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
sample_rate: int = 16000,
enable_speaker_diarization: bool = True,
enable_speaker_identification: bool = True,
enable_text_cleanup: bool = True,
word_timestamps: bool = False,
timestamp_scale: float = 1.0,
task_id: Optional[str] = None,
progress_callback: Optional[Callable[[str, str, int, Optional[dict[str, Any]]], None]] = None,
) -> ASRFullResult:
"""转录长音频文件(自动分段)
Args:
audio_path: 音频文件路径
hotwords: 热词
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
enable_speaker_diarization: 是否启用说话人分离
enable_speaker_identification: 是否提取声纹向量用于匹配已注册声纹库
enable_text_cleanup: 是否启用文本去重和口头语清理
word_timestamps: 是否返回字词级时间戳(仅部分模型支持)
timestamp_scale: Timestamp correction factor from audio normalization.
task_id: 任务ID(用于日志追踪)
Returns:
ASRFullResult: 包含完整文本、分段结果和时长的结果
"""
from app.utils.audio_splitter import AudioSplitter
# 开始性能计时
start_time = time.time()
start_perf = time.perf_counter()
model_id = getattr(self, 'model_id', 'unknown')
stage_timings_ms: Dict[str, float] = {
"duration_probe_ms": 0.0,
"speaker_diarization_ms": 0.0,
"vad_audio_split_ms": 0.0,
"asr_inference_ms": 0.0,
"speaker_embedding_ms": 0.0,
"cleanup_ms": 0.0,
"text_cleanup_ms": 0.0,
"hotword_resolver_ms": 0.0,
}
task_prefix = f"[{task_id}] " if task_id else ""
def emit_progress(
stage: str,
message: str,
percentage: int,
detail: Optional[dict[str, Any]] = None,
) -> None:
if progress_callback is None:
return
try:
progress_callback(stage, message, percentage, detail)
except Exception as exc:
logger.warning("%s进度回调失败: %s", task_prefix, exc)
logger.info(
f"{task_prefix}[transcribe_long_audio] 音频: {audio_path}, "
f"speaker_diarization={enable_speaker_diarization}, "
f"speaker_identification={enable_speaker_identification}, "
f"text_cleanup={enable_text_cleanup}, "
f"word_level={word_timestamps}"
)
try:
# 获取音频时长
duration_started = time.perf_counter()
duration = get_audio_duration(audio_path)
stage_timings_ms["duration_probe_ms"] = (
time.perf_counter() - duration_started
) * 1000
_log_asr_stage_timing(
"duration_probe",
stage_timings_ms["duration_probe_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
audio_path=audio_path,
)
logger.info(f"{task_prefix}[transcribe_long_audio] 音频时长: {duration:.2f}秒")
emit_progress(
"analyzing",
f"音频时长 {duration:.1f} 秒,正在进行分割。",
15,
{"audio_duration_seconds": round(duration, 2)},
)
# 统一使用分段处理
speaker_segments = None
audio_segments = None
if enable_speaker_diarization:
# 多说话人:使用说话人分离
from app.utils.speaker_diarizer import SpeakerDiarizer
logger.info(f"{task_prefix}使用说话人分离模式")
emit_progress(
"diarizing",
"正在执行说话人分离。",
18,
{"audio_duration_seconds": round(duration, 2)},
)
diarizer = SpeakerDiarizer()
diarization_started = time.perf_counter()
speaker_segments = diarizer.split_audio_by_speakers(audio_path)
stage_timings_ms["speaker_diarization_ms"] = (
time.perf_counter() - diarization_started
) * 1000
speaker_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in (speaker_segments or [])
)
_log_asr_stage_timing(
"speaker_diarization",
stage_timings_ms["speaker_diarization_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(speaker_segments or []),
segmented_audio_sec=round(speaker_audio_sec, 2),
)
if not speaker_segments:
logger.warning(f"{task_prefix}说话人分离未检测到片段,fallback 到 VAD 分割")
if not speaker_segments:
# 单说话人:使用 VAD 分割
logger.info(f"{task_prefix}使用 VAD 分割模式")
emit_progress(
"splitting",
"正在执行 VAD 音频分割。",
20,
{"audio_duration_seconds": round(duration, 2)},
)
splitter = AudioSplitter(device=self.device)
split_started = time.perf_counter()
audio_segments = splitter.split_audio_file(audio_path)
stage_timings_ms["vad_audio_split_ms"] = (
time.perf_counter() - split_started
) * 1000
split_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in (audio_segments or [])
)
_log_asr_stage_timing(
"vad_audio_split",
stage_timings_ms["vad_audio_split_ms"],
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(audio_segments or []),
segmented_audio_sec=round(split_audio_sec, 2),
)
# 选择要处理的片段
segments_to_process = speaker_segments if speaker_segments else audio_segments
if not segments_to_process:
raise DefaultServerErrorException("音频分割失败:未生成任何片段")
logger.info(f"{task_prefix}音频已分割为 {len(segments_to_process)} 段")
total_segments = len(segments_to_process)
emit_progress(
"transcribing",
"开始批量识别。",
30,
{
"segment_total": total_segments,
"segment_completed": 0,
"audio_duration_seconds": round(duration, 2),
},
)
results: List[ASRSegmentResult] = []
# 使用批处理推理
batch_size = settings.ASR_BATCH_SIZE
total_batches = (len(segments_to_process) + batch_size - 1) // batch_size
logger.info(
f"{task_prefix}使用批处理推理,batch_size={batch_size}, "
f"word_timestamps={word_timestamps}"
)
for batch_start in range(0, len(segments_to_process), batch_size):
batch_end = min(batch_start + batch_size, len(segments_to_process))
batch_segments = segments_to_process[batch_start:batch_end]
batch_index = batch_start // batch_size + 1
logger.info(
f"{task_prefix}推理批次 "
f"{batch_index}/{total_batches}: "
f"片段 {batch_start+1}-{batch_end}/{len(segments_to_process)}"
)
emit_progress(
"transcribing",
f"识别中:{batch_index}/{total_batches} 批,{batch_start}/{total_segments} 段",
min(88, 30 + int((batch_start / max(total_segments, 1)) * 58)),
{
"segment_total": total_segments,
"segment_completed": batch_start,
"batch_index": batch_index,
"batch_total": total_batches,
"audio_duration_seconds": round(duration, 2),
},
)
try:
# 批量推理,支持时间戳
batch_started = time.perf_counter()
batch_results = self._transcribe_batch(
segments=batch_segments,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
word_timestamps=word_timestamps,
)
batch_inference_ms = (time.perf_counter() - batch_started) * 1000
stage_timings_ms["asr_inference_ms"] += batch_inference_ms
batch_audio_sec = sum(
float(getattr(seg, "duration_sec", 0.0))
for seg in batch_segments
)
valid_batch_results = 0
batch_embedding_ms = 0.0
for seg, result in zip(batch_segments, batch_results):
if result and result.text:
valid_batch_results += 1
start_sec = float(getattr(seg, "start_sec", 0.0))
end_sec = float(getattr(seg, "end_sec", start_sec))
speaker_embedding = None
if (
speaker_segments
and settings.SPEAKER_DB_ENABLED
and enable_speaker_identification
):
try:
from app.core.database import pg_speaker_db
from app.services.speaker_registry import (
get_speaker_registry_service,
)
audio_data = getattr(seg, "audio_data", None)
if pg_speaker_db.is_connected and audio_data is not None:
embedding_started = time.perf_counter()
speaker_embedding = (
get_speaker_registry_service()
.extract_embedding_from_audio(audio_data)
)
batch_embedding_ms += (
time.perf_counter() - embedding_started
) * 1000
except Exception as exc:
logger.warning(
"%s提取片段声纹向量失败,保留原说话人编号: %s",
task_prefix,
exc,
)
results.append(
ASRSegmentResult(
text=result.text,
start_time=start_sec,
end_time=end_sec,
speaker_id=getattr(seg, "speaker_id", None),
speaker_embedding=speaker_embedding,
word_tokens=result.word_tokens if word_timestamps else None,
)
)
stage_timings_ms["speaker_embedding_ms"] += batch_embedding_ms
_log_asr_stage_timing(
"asr_batch",
batch_inference_ms,
task_id=task_id,
model_id=model_id,
audio_duration_sec=batch_audio_sec,
batch_index=batch_index,
batch_total=total_batches,
segment_start=batch_start + 1,
segment_end=batch_end,
segment_total=total_segments,
batch_segment_count=len(batch_segments),
valid_segment_count=valid_batch_results,
speaker_embedding_ms=round(batch_embedding_ms, 2),
)
logger.info(
f"{task_prefix}批次推理完成,有效片段: "
f"{len([r for r in batch_results if r and r.text])}"
)
emit_progress(
"transcribing",
f"识别中:{batch_index}/{total_batches} 批,{batch_end}/{total_segments} 段",
min(90, 30 + int((batch_end / max(total_segments, 1)) * 60)),
{
"segment_total": total_segments,
"segment_completed": batch_end,
"batch_index": batch_index,
"batch_total": total_batches,
"audio_duration_seconds": round(duration, 2),
},
)
except Exception as e:
logger.error(f"{task_prefix}批次推理失败: {e}, 跳过该批次")
emit_progress(
"transcribing",
f"第 {batch_index}/{total_batches} 批识别失败,已跳过。",
min(90, 30 + int((batch_end / max(total_segments, 1)) * 60)),
{
"segment_total": total_segments,
"segment_completed": batch_end,
"batch_index": batch_index,
"batch_total": total_batches,
"error": str(e),
},
)
# 清理临时文件(独立清理,避免条件遗漏)
try:
cleanup_started = time.perf_counter()
if speaker_segments:
from app.utils.speaker_diarizer import SpeakerDiarizer
SpeakerDiarizer.cleanup_segments(speaker_segments)
if audio_segments:
AudioSplitter.cleanup_segments(audio_segments)
stage_timings_ms["cleanup_ms"] = (
time.perf_counter() - cleanup_started
) * 1000
_log_asr_stage_timing(
"temp_cleanup",
stage_timings_ms["cleanup_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(segments_to_process),
)
except Exception as e:
logger.warning(f"清理临时文件时出错: {e}")
overlap_trimmed_count = 0
if enable_text_cleanup:
text_cleanup_started = time.perf_counter()
cleaned_results: List[ASRSegmentResult] = []
previous_text = ""
for seg in results:
cleaned_text = deduplicate_asr_text(seg.text)
cleaned_text, was_trimmed = trim_segment_boundary_overlap(
previous_text,
cleaned_text,
)
if was_trimmed:
overlap_trimmed_count += 1
if not cleaned_text:
continue
seg.text = cleaned_text
cleaned_results.append(seg)
previous_text = cleaned_text
if len(cleaned_results) != len(results) or overlap_trimmed_count:
logger.info(
"%s文字去重完成:%s -> %s 段,边界重叠裁剪 %s 次",
task_prefix,
len(results),
len(cleaned_results),
overlap_trimmed_count,
)
results = cleaned_results
stage_timings_ms["text_cleanup_ms"] = (
time.perf_counter() - text_cleanup_started
) * 1000
_log_asr_stage_timing(
"text_cleanup",
stage_timings_ms["text_cleanup_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(results),
overlap_trimmed_count=overlap_trimmed_count,
cleanup_enabled=enable_text_cleanup,
)
hotword_resolver_started = time.perf_counter()
hotword_resolver_applied_count = 0
matched_hotwords: list[str] = []
if hotwords:
for seg in results:
resolved_text, was_applied, matched = apply_hotword_rules(
seg.text,
hotwords,
)
if resolved_text:
seg.text = resolved_text
if was_applied:
hotword_resolver_applied_count += 1
for hotword_text in matched:
if hotword_text not in matched_hotwords:
matched_hotwords.append(hotword_text)
if hotwords:
stage_timings_ms["hotword_resolver_ms"] = (
time.perf_counter() - hotword_resolver_started
) * 1000
_log_asr_stage_timing(
"hotword_resolver",
stage_timings_ms["hotword_resolver_ms"],
task_id=task_id,
model_id=model_id,
segment_count=len(results),
applied_segment_count=hotword_resolver_applied_count,
matched_hotwords=matched_hotwords,
)
if hotword_resolver_applied_count:
logger.info(
"%s热词纠偏完成:命中 %s 段,热词=%s",
task_prefix,
hotword_resolver_applied_count,
matched_hotwords,
)
all_texts = [seg.text for seg in results]
full_text = "\n".join(all_texts)
emit_progress(
"finalizing",
"分段识别完成,正在整理文本。",
92,
{
"segment_total": len(segments_to_process),
"segment_completed": len(segments_to_process),
"valid_segment_count": len(results),
"text_cleanup_enabled": enable_text_cleanup,
"overlap_trimmed_count": overlap_trimmed_count,
"audio_duration_seconds": round(duration, 2),
},
)
logger.info(
f"长音频识别完成,共 {len(results)} 个有效分段,"
f"总字符数: {len(full_text)}"
)
# 计算性能指标
total_duration_ms = (time.time() - start_time) * 1000
total_perf_ms = (time.perf_counter() - start_perf) * 1000
if timestamp_scale != 1.0:
for seg in results:
seg.start_time *= timestamp_scale
seg.end_time *= timestamp_scale
if seg.word_tokens:
for word_token in seg.word_tokens:
word_token.start_time *= timestamp_scale
word_token.end_time *= timestamp_scale
duration *= timestamp_scale
logger.info(
f"{task_prefix}Timestamp scaling applied: scale={timestamp_scale:.6f}"
)
_log_asr_stage_timing(
"asr_total",
total_perf_ms,
task_id=task_id,
model_id=model_id,
audio_duration_sec=duration,
segment_count=len(results),
total_segment_count=len(segments_to_process),
**{key: round(value, 2) for key, value in stage_timings_ms.items()},
)
log_inference_metrics(
logger=logger,
message="长音频识别完成",
task_id=task_id,
duration_ms=total_duration_ms,
audio_duration_sec=duration,
model_id=model_id,
status="success",
segments_count=len(results),
batch_size=settings.ASR_BATCH_SIZE,
event="asr_inference_metrics",
**{key: round(value, 2) for key, value in stage_timings_ms.items()},
enable_speaker_diarization=enable_speaker_diarization,
enable_speaker_identification=enable_speaker_identification,
enable_text_cleanup=enable_text_cleanup,
word_timestamps=word_timestamps,
)
return ASRFullResult(
text=full_text,
segments=results,
duration=duration,
)
except Exception as e:
# 计算失败时的性能指标
total_duration_ms = (time.time() - start_time) * 1000
try:
duration = get_audio_duration(audio_path)
except Exception:
duration = 0
log_inference_metrics(
logger=logger,
message="长音频识别失败",
task_id=task_id,
duration_ms=total_duration_ms,
audio_duration_sec=duration,
model_id=model_id,
status="error",
error=str(e),
)
logger.error(f"长音频识别失败: {e}")
raise DefaultServerErrorException(f"长音频识别失败: {str(e)}")
@abstractmethod
def is_model_loaded(self) -> bool:
"""检查模型是否已加载"""
pass
@property
@abstractmethod
def device(self) -> str:
"""获取设备信息"""
pass
@property
@abstractmethod
def supports_realtime(self) -> bool:
"""是否支持实时识别"""
pass
def _transcribe_batch(
self,
segments: List[Any],
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
sample_rate: int = 16000,
word_timestamps: bool = False,
) -> List[ASRSegmentResult]:
"""批量推理多个音频片段
Args:
segments: 音频片段列表(每个片段需要有 temp_file 属性)
hotwords: 热词
enable_punctuation: 是否启用标点
enable_itn: 是否启用 ITN
sample_rate: 采样率
word_timestamps: 是否返回字词级时间戳
Returns:
ASRSegmentResult 列表,与输入片段一一对应
"""
# 默认实现:逐个推理(子类可以重写实现真正的批处理)
results = []
for idx, seg in enumerate(segments):
try:
if not seg.temp_file:
logger.warning(f"批处理片段 {idx + 1} 临时文件不存在,跳过")
results.append(ASRSegmentResult(text="", start_time=0.0, end_time=0.0))
continue
if word_timestamps:
# 需要时间戳:使用 transcribe_file_with_vad
raw_result = self.transcribe_file_with_vad(
audio_path=seg.temp_file,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
word_timestamps=True,
)
if raw_result.segments:
result_seg = raw_result.segments[0]
results.append(
ASRSegmentResult(
text=result_seg.text,
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
word_tokens=result_seg.word_tokens,
)
)
else:
results.append(
ASRSegmentResult(
text=raw_result.text,
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
)
)
else:
# 不需要时间戳:使用 transcribe_file
text = self.transcribe_file(
audio_path=seg.temp_file,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
enable_vad=False,
sample_rate=sample_rate,
)
results.append(
ASRSegmentResult(
text=text or "",
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, 'speaker_id', None),
)
)
except Exception as e:
logger.error(f"批处理片段 {idx + 1} 推理失败: {e}")
results.append(
ASRSegmentResult(
text="",
start_time=getattr(seg, 'start_sec', 0.0),
end_time=getattr(seg, 'end_sec', 0.0),
speaker_id=getattr(seg, 'speaker_id', None),
)
)
return results
@staticmethod
def _detect_device(device: str = "auto") -> str:
"""检测可用设备"""
from app.core.device import detect_device
return detect_device(device)
class RealTimeASREngine(BaseASREngine):
"""实时ASR引擎抽象基类"""
@property
def supports_realtime(self) -> bool:
"""支持实时识别"""
return True
@abstractmethod
def transcribe_websocket(
self,
audio_chunk: bytes,
cache: Optional[Dict] = None,
is_final: bool = False,
**kwargs,
) -> str:
"""WebSocket流式语音识别"""
pass

View File

@ -0,0 +1,141 @@
# -*- coding: utf-8 -*-
"""
全局VAD/PUNC模型管理模块
提供线程安全的全局模型实例管理
"""
import logging
import threading
from funasr import AutoModel
from app.core.config import settings
from app.infrastructure import resolve_model_path
logger = logging.getLogger(__name__)
# 全局语音活动检测(VAD)模型缓存(避免重复加载)
_global_vad_model = None
_vad_model_lock = threading.Lock()
_vad_inference_lock = threading.Lock() # 推理互斥锁,防止并发状态混乱
# 全局标点符号模型缓存(避免重复加载)
_global_punc_model = None
_punc_model_lock = threading.Lock()
_punc_inference_lock = threading.Lock() # 推理互斥锁,防止并发状态混乱
# 全局实时标点符号模型缓存(避免重复加载)
_global_punc_realtime_model = None
_punc_realtime_model_lock = threading.Lock()
_punc_realtime_inference_lock = threading.Lock() # 推理互斥锁,防止并发状态混乱
def _resolve_device(device: str) -> str:
"""解析设备字符串,将 auto 转换为实际的设备"""
from app.core.device import detect_device
return detect_device(device)
def get_global_vad_model(device: str):
"""获取全局语音活动检测(VAD)模型实例(线程安全,双重检查锁定)"""
global _global_vad_model
if _global_vad_model is None:
with _vad_model_lock:
if _global_vad_model is None:
try:
# 解析模型路径:优先使用本地缓存
resolved_vad_path = resolve_model_path(settings.VAD_MODEL)
logger.info(f"正在加载全局语音活动检测(VAD)模型: {resolved_vad_path}")
# 解析 auto 设备
resolved_device = _resolve_device(device)
_global_vad_model = AutoModel(
model=resolved_vad_path,
device=resolved_device,
speech_noise_thres=0.6, # VAD 语音噪声阈值(FunASR默认0.6,设为0.7稍微严格一些,分段更碎)
**settings.FUNASR_AUTOMODEL_KWARGS,
)
logger.info("全局语音活动检测(VAD)模型加载成功 (speech_noise_thres=0.6)")
except Exception as e:
logger.error(f"全局语音活动检测(VAD)模型加载失败: {str(e)}")
_global_vad_model = None
raise
return _global_vad_model
def get_vad_inference_lock():
"""获取VAD模型推理锁(线程安全)"""
return _vad_inference_lock
def get_global_punc_model(device: str):
"""获取全局标点符号模型实例(离线版,线程安全,双重检查锁定)"""
global _global_punc_model
if _global_punc_model is None:
with _punc_model_lock:
if _global_punc_model is None:
try:
# 解析模型路径:优先使用本地缓存
resolved_punc_path = resolve_model_path(settings.PUNC_MODEL)
logger.info(f"正在加载全局标点符号模型(离线): {resolved_punc_path}")
# 解析 auto 设备
resolved_device = _resolve_device(device)
_global_punc_model = AutoModel(
model=resolved_punc_path,
device=resolved_device,
**settings.FUNASR_AUTOMODEL_KWARGS,
)
logger.info("全局标点符号模型(离线)加载成功")
except Exception as e:
logger.error(f"全局标点符号模型(离线)加载失败: {str(e)}")
_global_punc_model = None
raise
return _global_punc_model
def get_punc_inference_lock():
"""获取PUNC模型推理锁(线程安全)"""
return _punc_inference_lock
def get_global_punc_realtime_model(device: str):
"""获取全局实时标点符号模型实例(线程安全,双重检查锁定)"""
global _global_punc_realtime_model
if _global_punc_realtime_model is None:
with _punc_realtime_model_lock:
if _global_punc_realtime_model is None:
try:
# 解析模型路径:优先使用本地缓存
resolved_punc_realtime_path = resolve_model_path(settings.PUNC_REALTIME_MODEL)
logger.info(f"正在加载全局标点符号模型(实时): {resolved_punc_realtime_path}")
# 解析 auto 设备
resolved_device = _resolve_device(device)
_global_punc_realtime_model = AutoModel(
model=resolved_punc_realtime_path,
device=resolved_device,
**settings.FUNASR_AUTOMODEL_KWARGS,
)
logger.info("全局标点符号模型(实时)加载成功")
except Exception as e:
logger.error(f"全局标点符号模型(实时)加载失败: {str(e)}")
_global_punc_realtime_model = None
raise
return _global_punc_realtime_model
def get_punc_realtime_inference_lock():
"""获取实时PUNC模型推理锁(线程安全)"""
return _punc_realtime_inference_lock

View File

@ -0,0 +1,5 @@
# -*- coding: utf-8 -*-
"""
模型实现模块
包含需要本地代码的自定义模型实现
"""

View File

@ -0,0 +1,215 @@
# -*- coding: utf-8 -*-
"""ASR model metadata and engine factory."""
import json
import threading
import logging
from typing import Dict, Any, Optional, List
from pathlib import Path
from typing import Callable
from ...core.config import settings
from ...core.exceptions import DefaultServerErrorException, InvalidParameterException
from .engines import BaseASREngine
from .model_plan import get_default_model_id
logger = logging.getLogger(__name__)
# 引擎注册表(使用Any避免循环导入问题)
_ENGINE_REGISTRY: Dict[str, Callable[[Any], BaseASREngine]] = {}
def _supports_qwen_realtime_on_device(configured_device: str) -> bool:
"""Resolve whether Qwen realtime mode is available on the active device."""
from app.core.device import detect_device
from app.core.accelerator import get_accelerator_info
from .qwenasr_rust import is_qwenasr_rust_available
device = detect_device(configured_device)
accelerator = get_accelerator_info()
if accelerator.is_gpu and device.startswith("cuda"):
return True
if device == "cpu":
return is_qwenasr_rust_available()
return False
def register_engine(engine_type: str, factory: Callable[[Any], BaseASREngine]):
"""注册ASR引擎工厂函数"""
_ENGINE_REGISTRY[engine_type] = factory
logger.info(f"注册引擎类型: {engine_type}")
class DeclaredEntryConfig:
"""声明条目配置,可表示模型或 capability。"""
def __init__(self, model_id: str, config: Dict[str, Any]):
self.model_id = model_id
self.name = config["name"]
self.kind = config.get("kind", "model")
self.engine = config["engine"]
self.description = config.get("description", "")
self.languages = config.get("languages", [])
self.supports_realtime = config.get("supports_realtime", False)
# 模型路径结构
self.models = config.get("models", {})
self.offline_model_path = self.models.get("offline")
self.realtime_model_path = self.models.get("realtime")
# 额外参数(如 trust_remote_code 等)
self.extra_kwargs = config.get("extra_kwargs", {})
@property
def has_offline_model(self) -> bool:
"""是否有离线模型"""
return bool(self.offline_model_path)
@property
def has_realtime_model(self) -> bool:
"""是否有实时模型"""
return bool(self.realtime_model_path)
class ModelManager:
"""Static model metadata plus engine construction."""
def __init__(self):
self._declared_entry_configs: Dict[str, DeclaredEntryConfig] = {}
self._default_model_id: Optional[str] = None
self._load_models_config()
def _load_models_config(self) -> None:
"""加载模型配置文件"""
models_file = Path(settings.models_config_path)
if not models_file.exists():
raise DefaultServerErrorException("models.json 配置文件不存在")
try:
with open(models_file, "r", encoding="utf-8") as f:
config = json.load(f)
for model_id, model_config in config["models"].items():
self._declared_entry_configs[model_id] = DeclaredEntryConfig(model_id, model_config)
self._default_model_id = get_default_model_id(
all_model_ids=list(self._declared_entry_configs.keys()),
)
if not self._default_model_id and self._declared_entry_configs:
self._default_model_id = list(self._declared_entry_configs.keys())[0]
except (json.JSONDecodeError, KeyError) as e:
raise DefaultServerErrorException(f"模型配置文件格式错误: {str(e)}")
def get_declared_entry_config(self, model_id: Optional[str] = None) -> DeclaredEntryConfig:
"""获取声明条目配置。"""
if model_id is None:
model_id = self._default_model_id
if not model_id:
raise InvalidParameterException("未指定模型且没有默认模型")
if model_id not in self._declared_entry_configs:
available_models = ", ".join(self._declared_entry_configs.keys())
raise InvalidParameterException(
f"未知的模型: {model_id},可用模型: {available_models}"
)
return self._declared_entry_configs[model_id]
def list_declared_entries(self) -> List[Dict[str, Any]]:
"""列出声明的模型与 capability 元数据。"""
entries = []
for model_id, config in self._declared_entry_configs.items():
offline_path_exists = False
realtime_path_exists = False
if config.offline_model_path:
offline_model_path = (
Path(settings.MODELSCOPE_PATH) / config.offline_model_path
)
offline_path_exists = offline_model_path.exists()
if config.realtime_model_path:
realtime_model_path = (
Path(settings.MODELSCOPE_PATH) / config.realtime_model_path
)
realtime_path_exists = realtime_model_path.exists()
supports_realtime = config.supports_realtime
if config.engine == "qwen3":
supports_realtime = _supports_qwen_realtime_on_device(settings.DEVICE)
entries.append(
{
"id": model_id,
"kind": config.kind,
"name": config.name,
"engine": config.engine,
"description": config.description,
"languages": config.languages,
"default": model_id == self._default_model_id,
"supports_realtime": supports_realtime,
"offline_model": (
{
"path": config.offline_model_path,
"exists": offline_path_exists,
}
if config.offline_model_path
else None
),
"realtime_model": (
{
"path": config.realtime_model_path,
"exists": realtime_path_exists,
}
if config.realtime_model_path
else None
),
}
)
return entries
def _create_engine(self, config: DeclaredEntryConfig) -> BaseASREngine:
"""创建ASR引擎实例"""
engine_type = config.engine.lower()
factory = _ENGINE_REGISTRY.get(engine_type)
if not factory:
raise InvalidParameterException(
f"不支持的引擎类型: {config.engine}"
)
return factory(config)
def create_engine(self, model_id: Optional[str] = None) -> BaseASREngine:
"""Create a fresh engine instance."""
config = self.get_declared_entry_config(model_id)
return self._create_engine(config)
# 全局模型管理器实例
_model_manager: Optional[ModelManager] = None
_model_manager_lock = threading.Lock()
def get_model_manager() -> ModelManager:
"""获取全局模型管理器实例(线程安全)"""
global _model_manager
if _model_manager is None:
with _model_manager_lock:
if _model_manager is None:
_model_manager = ModelManager()
return _model_manager
# 注册内置引擎
def _register_builtin_engines():
"""注册内置的ASR引擎"""
try:
from .qwen3_engine import Qwen3ASREngine # noqa: F401
from .qwen3_engine import _register_qwen3_engine
_register_qwen3_engine(register_engine, DeclaredEntryConfig)
except ImportError as e:
logger.warning(f"Qwen3引擎不可用: {e}")
# 模块加载时自动注册内置引擎
_register_builtin_engines()

View File

@ -0,0 +1,230 @@
# -*- coding: utf-8 -*-
"""Shared capability-to-model asset definitions."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal, Optional
from app.core.config import settings
from app.services.asr.manager import get_model_manager
from app.services.asr.model_plan import (
get_active_qwen_model,
load_supported_model_ids,
)
ModelSource = Literal["modelscope"]
@dataclass(frozen=True)
class ModelAsset:
source: ModelSource
model_id: str
description: str
revision: Optional[str] = None
required_patterns: tuple[str, ...] = ()
alternative_required_patterns: tuple[tuple[str, ...], ...] = ()
min_total_size_bytes: int = 0
_VAD_ASSETS = (
ModelAsset(
source="modelscope",
model_id=settings.VAD_MODEL,
description="VAD",
revision="v2.0.2",
required_patterns=("configuration.json", "config.yaml", "model.pb"),
min_total_size_bytes=1_000_000,
),
)
_DIARIZATION_ASSETS = (
ModelAsset(
source="modelscope",
model_id="iic/speech_campplus_speaker-diarization_common",
description="CAM++ Diarization",
required_patterns=(
"configuration.json",
"config.yaml",
"onnx/asd.onnx",
"onnx/face_recog_ir101.onnx",
"onnx/fqa.onnx",
"onnx/version-RFB-320.onnx",
),
min_total_size_bytes=50_000_000,
),
ModelAsset(
source="modelscope",
model_id=settings.SV_MODEL,
description="Configured Speaker Verification",
required_patterns=("configuration.json", "config.yaml", "campplus_cn_common.bin"),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id=settings.REALTIME_SV_MODEL,
description="Realtime Speaker Verification",
required_patterns=("configuration.json",),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id="damo/speech_campplus_sv_zh-cn_16k-common",
description="CAM++ Speaker Verification",
required_patterns=("configuration.json", "config.yaml", "campplus_cn_common.bin"),
min_total_size_bytes=10_000_000,
),
ModelAsset(
source="modelscope",
model_id="damo/speech_campplus-transformer_scl_zh-cn_16k-common",
description="CAM++ Transformer",
required_patterns=("configuration.json", "campplus_cn_encoder.pt", "transformer_backend.pt"),
min_total_size_bytes=10_000_000,
),
)
def get_download_modelscope_assets() -> list[ModelAsset]:
"""Return the full static ModelScope export set used by predownload/export."""
return _dedupe_assets([
*_VAD_ASSETS,
*_DIARIZATION_ASSETS,
])
def get_runtime_required_modelscope_assets(
*,
include_realtime_punc: bool,
) -> list[ModelAsset]:
"""Return ModelScope assets required by the current runtime plan."""
_ = include_realtime_punc
return _dedupe_assets([*_VAD_ASSETS, *_DIARIZATION_ASSETS])
def _dedupe_assets(assets: list[ModelAsset]) -> list[ModelAsset]:
deduped: list[ModelAsset] = []
seen: set[tuple[str, str]] = set()
for asset in assets:
key = (asset.source, asset.model_id)
if key in seen:
continue
seen.add(key)
deduped.append(asset)
return deduped
def get_camplusplus_replacement_paths(cache_dir: str) -> dict[str, str]:
"""Return the CAM++ offline replacement map for local cache paths."""
return {
"damo/speech_campplus_sv_zh-cn_16k-common": f"{cache_dir}/damo/speech_campplus_sv_zh-cn_16k-common",
"iic/speech_campplus_sv_zh-cn_16k-common": f"{cache_dir}/iic/speech_campplus_sv_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": f"{cache_dir}/damo/speech_campplus-transformer_scl_zh-cn_16k-common",
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": f"{cache_dir}/damo/speech_campplus-transformer_scl_zh-cn-16k-common",
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": f"{cache_dir}/damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
}
_QWEN_MODELSCOPE_MODEL_IDS = {
"Qwen/Qwen3-ASR-0.6B": "Qwen/Qwen3-ASR-0.6B",
"Qwen/Qwen3-ASR-1.7B": "Qwen/Qwen3-ASR-1.7B",
"Qwen/Qwen3-ForcedAligner-0.6B": "Qwen/Qwen3-ForcedAligner-0.6B",
}
def get_qwen_modelscope_model_id(model_id: str) -> Optional[str]:
"""Map runtime Qwen model ID to the corresponding ModelScope model ID.
Returns:
ModelScope model ID if available, None otherwise.
"""
return _QWEN_MODELSCOPE_MODEL_IDS.get(model_id)
def get_enabled_qwen_modelscope_assets(
*,
include_forced_aligner: bool = True,
) -> list[ModelAsset]:
"""Return ModelScope assets required by the runtime Qwen plan."""
manager = get_model_manager()
assets: list[ModelAsset] = []
model_id = get_active_qwen_model()
model_config = manager.get_declared_entry_config(model_id)
offline_model = model_config.offline_model_path
if offline_model:
ms_model_id = get_qwen_modelscope_model_id(offline_model)
if ms_model_id:
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_model_id,
description=f"{model_config.name} Offline (ModelScope)",
required_patterns=("config.json",),
alternative_required_patterns=(
("model.safetensors",),
("model-*.safetensors",),
),
min_total_size_bytes=500_000_000,
)
)
forced_aligner = str(model_config.extra_kwargs.get("forced_aligner_path") or "").strip()
if forced_aligner and include_forced_aligner:
ms_aligner_id = get_qwen_modelscope_model_id(forced_aligner)
if ms_aligner_id:
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_aligner_id,
description=f"{model_config.name} Forced Aligner (ModelScope)",
required_patterns=("config.json", "model.safetensors"),
min_total_size_bytes=500_000_000,
)
)
return assets
def get_all_qwen_modelscope_assets(
*,
include_forced_aligner: bool = True,
) -> list[ModelAsset]:
"""Return all declared Qwen ModelScope assets for offline bundles."""
manager = get_model_manager()
assets: list[ModelAsset] = []
seen_model_ids: set[str] = set()
for model_id in sorted(load_supported_model_ids()):
if not model_id.startswith("qwen3-asr-"):
continue
model_config = manager.get_declared_entry_config(model_id)
offline_model = model_config.offline_model_path
if offline_model:
ms_model_id = get_qwen_modelscope_model_id(offline_model)
if ms_model_id and ms_model_id not in seen_model_ids:
seen_model_ids.add(ms_model_id)
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_model_id,
description=f"{model_config.name} Offline (ModelScope)",
required_patterns=("config.json",),
alternative_required_patterns=(
("model.safetensors",),
("model-*.safetensors",),
),
min_total_size_bytes=500_000_000,
)
)
forced_aligner = str(model_config.extra_kwargs.get("forced_aligner_path") or "").strip()
if forced_aligner and include_forced_aligner:
ms_aligner_id = get_qwen_modelscope_model_id(forced_aligner)
if ms_aligner_id and ms_aligner_id not in seen_model_ids:
seen_model_ids.add(ms_aligner_id)
assets.append(
ModelAsset(
source="modelscope",
model_id=ms_aligner_id,
description="Qwen3 Forced Aligner (ModelScope)",
required_patterns=("config.json", "model.safetensors"),
min_total_size_bytes=500_000_000,
)
)
return assets

View File

@ -0,0 +1,105 @@
# -*- coding: utf-8 -*-
"""Single-source deployment model planning."""
from __future__ import annotations
import json
import os
import platform
from pathlib import Path
from typing import Optional
from app.core.config import settings
QWEN_MODEL_OVERRIDE_ENV = "QWEN3_ASR_MODEL"
_QWEN_MODEL_ALIASES = {
"qwen3-asr-0.6b": "qwen3-asr-0.6b",
"0.6b": "qwen3-asr-0.6b",
"0.6": "qwen3-asr-0.6b",
"qwen/qwen3-asr-0.6b": "qwen3-asr-0.6b",
"qwen3-asr-1.7b": "qwen3-asr-1.7b",
"1.7b": "qwen3-asr-1.7b",
"1.7": "qwen3-asr-1.7b",
"qwen/qwen3-asr-1.7b": "qwen3-asr-1.7b",
}
def load_supported_model_ids() -> list[str]:
"""Load declared model ids from models.json."""
models_file = Path(settings.models_config_path)
if not models_file.exists():
return []
with open(models_file, "r", encoding="utf-8") as f:
config = json.load(f)
return list(config.get("models", {}).keys())
def get_qwen_model_override() -> Optional[str]:
"""Return the explicit Qwen model override from the environment."""
raw_value = (os.getenv(QWEN_MODEL_OVERRIDE_ENV) or "").strip()
if not raw_value:
return None
normalized = _QWEN_MODEL_ALIASES.get(raw_value.lower(), raw_value)
return normalized
def detect_qwen_model_by_vram(all_model_ids: Optional[list[str]] = None) -> Optional[str]:
"""Pick the active Qwen model for the current machine."""
from app.core.accelerator import get_accelerator_info
from app.core.device import detect_device, get_vram_gb
from app.services.asr.qwenasr_rust import is_qwenasr_rust_available
model_ids = all_model_ids or load_supported_model_ids()
override_model = get_qwen_model_override()
if override_model:
return override_model if override_model in model_ids else None
resolved_device = detect_device(settings.DEVICE)
accelerator = get_accelerator_info()
# macOS defaults to the lighter Rust CPU path unless QWEN3_ASR_MODEL is set.
if platform.system() == "Darwin":
return "qwen3-asr-0.6b" if is_qwenasr_rust_available() and "qwen3-asr-0.6b" in model_ids else None
if resolved_device == "cpu" or not accelerator.is_gpu:
return "qwen3-asr-0.6b" if is_qwenasr_rust_available() and "qwen3-asr-0.6b" in model_ids else None
vram = get_vram_gb()
preferred = "qwen3-asr-1.7b" if vram >= 32 else "qwen3-asr-0.6b"
if preferred in model_ids:
return preferred
fallback = "qwen3-asr-0.6b" if preferred == "qwen3-asr-1.7b" else "qwen3-asr-1.7b"
return fallback if fallback in model_ids else None
def get_active_qwen_model(all_model_ids: Optional[list[str]] = None) -> str:
"""Return the required Qwen model for the current machine."""
model_ids = all_model_ids or load_supported_model_ids()
qwen_model = detect_qwen_model_by_vram(model_ids)
if not qwen_model:
override_model = get_qwen_model_override()
if override_model:
available_qwen_models = ", ".join(
model_id for model_id in model_ids if model_id.startswith("qwen")
)
raise RuntimeError(
f"{QWEN_MODEL_OVERRIDE_ENV}={override_model} 不在可用 Qwen3-ASR 模型中: "
f"{available_qwen_models}"
)
raise RuntimeError("当前环境未找到可运行的 Qwen3-ASR 模型")
return qwen_model
def get_runtime_model_ids(all_model_ids: Optional[list[str]] = None) -> list[str]:
"""Return the runtime model/capability plan for the current machine."""
model_ids = all_model_ids or load_supported_model_ids()
return [get_active_qwen_model(model_ids)]
def get_default_model_id(all_model_ids: Optional[list[str]] = None) -> str:
"""Return the single default offline model for API/UI selection."""
return get_active_qwen_model(all_model_ids)

View File

@ -0,0 +1,93 @@
# -*- coding: utf-8 -*-
"""Offline/realtime model selection helpers."""
from __future__ import annotations
from typing import List, Optional
from ...core.exceptions import InvalidParameterException
from .manager import get_model_manager
from .model_plan import (
get_active_qwen_model,
get_default_model_id,
get_runtime_model_ids,
)
def get_active_qwen_model_id() -> str:
"""Return the currently active Qwen model id."""
active_qwen_model = get_active_qwen_model()
return active_qwen_model or "qwen3-asr-0.6b"
def get_offline_model_ids() -> List[str]:
"""Return enabled offline-capable models for docs and APIs."""
manager = get_model_manager()
runtime_models = get_runtime_model_ids()
def sort_key(model_id: str) -> tuple[int, str]:
if model_id.startswith("qwen"):
return (0, model_id)
return (1, model_id)
offline_models = [
model_id
for model_id in runtime_models
if manager.get_declared_entry_config(model_id).has_offline_model
]
return sorted(offline_models, key=sort_key)
def get_default_offline_model_id() -> str:
"""Return the default offline-capable model."""
default_model = get_default_model_id()
if default_model:
try:
if get_model_manager().get_declared_entry_config(default_model).has_offline_model:
return default_model
except InvalidParameterException:
pass
return get_active_qwen_model_id()
def validate_offline_model_id(model_id: Optional[str]) -> str:
"""Validate offline-capable model ids for REST transcription requests."""
available_models = get_offline_model_ids()
if not model_id or not model_id.strip():
return get_default_offline_model_id()
requested_model = model_id.strip()
if requested_model.lower() == "qwen3-asr":
active_qwen_model = get_active_qwen_model_id()
if active_qwen_model.startswith("qwen") and active_qwen_model in available_models:
return active_qwen_model
raise InvalidParameterException("当前环境未启用 Qwen3-ASR 模型")
if requested_model not in available_models:
raise InvalidParameterException(
f"不支持的离线模型ID: {requested_model}。可用模型: {', '.join(available_models)}"
)
return requested_model
def validate_realtime_model_id(model_id: Optional[str]) -> str:
"""Validate realtime-capable model ids for websocket protocols."""
available_models = get_offline_model_ids()
if not model_id:
return get_default_offline_model_id()
if model_id.lower() == "qwen3-asr":
active_qwen_model = get_active_qwen_model_id()
if active_qwen_model.startswith("qwen") and active_qwen_model in available_models:
return active_qwen_model
raise InvalidParameterException("当前环境未启用 Qwen3-ASR 模型")
if model_id not in available_models:
raise InvalidParameterException(
f"不支持的模型ID: {model_id}。可用模型: {', '.join(available_models)}"
)
return model_id

View File

@ -0,0 +1,70 @@
{
"models": {
"qwen3-asr-1.7b": {
"name": "Qwen3-ASR-1.7B",
"kind": "model",
"engine": "qwen3",
"description": "Qwen3-ASR 1.7B,支持52种语言和方言;CUDA 使用 vLLM,CPU/macOS 使用 QwenASR Rust backend",
"languages": [
"zh",
"en",
"yue",
"ja",
"ko",
"ar",
"de",
"es",
"fr",
"pt",
"id",
"it",
"ru",
"th",
"vi"
],
"default": true,
"supports_realtime": true,
"models": {
"offline": "Qwen/Qwen3-ASR-1.7B"
},
"extra_kwargs": {
"max_model_len": 16384,
"forced_aligner_path": "Qwen/Qwen3-ForcedAligner-0.6B",
"max_inference_batch_size": 16
}
},
"qwen3-asr-0.6b": {
"name": "Qwen3-ASR-0.6B",
"kind": "model",
"engine": "qwen3",
"description": "Qwen3-ASR 0.6B,轻量版支持52种语言和方言;CUDA 使用 vLLM,CPU/macOS 使用 QwenASR Rust backend",
"languages": [
"zh",
"en",
"yue",
"ja",
"ko",
"ar",
"de",
"es",
"fr",
"pt",
"id",
"it",
"ru",
"th",
"vi"
],
"default": false,
"supports_realtime": true,
"models": {
"offline": "Qwen/Qwen3-ASR-0.6B"
},
"extra_kwargs": {
"max_model_len": 16384,
"forced_aligner_path": "Qwen/Qwen3-ForcedAligner-0.6B",
"max_inference_batch_size": 16
}
}
}
}

View File

@ -0,0 +1,137 @@
# -*- coding: utf-8 -*-
"""Shared offline transcription workflow."""
from __future__ import annotations
from dataclasses import dataclass
import logging
from typing import Any, Callable, Optional
from fastapi import Request
from app.models.common import SampleRate
from app.services.asr.engines import ASRFullResult
from app.services.asr.model_selection import validate_offline_model_id
from app.services.asr.runtime import OfflineASRRequest, get_runtime_router
from app.services.audio import get_audio_service
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class PreparedAudio:
normalized_path: str
duration: float
original_path: str
timestamp_scale: float = 1.0
@dataclass(frozen=True)
class OfflineTranscriptionOptions:
model_id: Optional[str] = None
sample_rate: int = 16000
hotwords: str = ""
enable_speaker_diarization: bool = True
enable_speaker_identification: bool = True
enable_text_cleanup: bool = True
word_timestamps: bool = False
task_id: Optional[str] = None
progress_callback: Optional[
Callable[[str, str, int, Optional[dict[str, Any]]], None]
] = None
class OfflineTranscriptionService:
"""Prepare audio and run the active offline ASR model."""
def __init__(self) -> None:
self._audio_service = get_audio_service()
async def prepare_from_request(
self,
*,
request: Request,
audio_address: Optional[str],
task_id: str,
sample_rate: int,
) -> PreparedAudio:
audio = await self._audio_service.process_from_request(
request=request,
audio_address=audio_address,
task_id=task_id,
sample_rate=sample_rate,
)
return PreparedAudio(
normalized_path=audio.normalized_path,
duration=audio.duration,
original_path=audio.original_path,
timestamp_scale=audio.timestamp_scale,
)
async def prepare_upload(
self,
*,
audio_data: bytes,
filename: Optional[str],
task_id: str,
sample_rate: int,
) -> PreparedAudio:
audio = await self._audio_service.process_upload_file(
audio_data=audio_data,
filename=filename,
task_id=task_id,
sample_rate=sample_rate,
)
return PreparedAudio(
normalized_path=audio.normalized_path,
duration=audio.duration,
original_path=audio.original_path,
timestamp_scale=audio.timestamp_scale,
)
async def transcribe(
self,
prepared_audio: PreparedAudio,
options: OfflineTranscriptionOptions,
) -> ASRFullResult:
model_id = validate_offline_model_id(options.model_id)
logger.info(
"ASR model resolved: requested=%s, resolved=%s",
options.model_id,
model_id,
)
return await get_runtime_router().run_offline(
OfflineASRRequest(
model_id=model_id,
audio_path=prepared_audio.normalized_path,
hotwords=options.hotwords,
enable_punctuation=True,
enable_itn=True,
sample_rate=options.sample_rate or int(SampleRate.RATE_16000),
enable_speaker_diarization=options.enable_speaker_diarization,
enable_speaker_identification=options.enable_speaker_identification,
enable_text_cleanup=options.enable_text_cleanup,
word_timestamps=options.word_timestamps,
timestamp_scale=prepared_audio.timestamp_scale,
task_id=options.task_id,
progress_callback=options.progress_callback,
)
)
def cleanup(self, prepared_audio: Optional[PreparedAudio]) -> None:
if prepared_audio is None:
return
self._audio_service.cleanup(
prepared_audio.original_path,
prepared_audio.normalized_path,
)
_offline_transcription_service: Optional[OfflineTranscriptionService] = None
def get_offline_transcription_service() -> OfflineTranscriptionService:
global _offline_transcription_service
if _offline_transcription_service is None:
_offline_transcription_service = OfflineTranscriptionService()
return _offline_transcription_service

View File

@ -0,0 +1,693 @@
# -*- coding: utf-8 -*-
"""Qwen3-ASR engine with official vLLM and vendored Rust backends."""
import logging
import os
from concurrent.futures import ThreadPoolExecutor
from typing import Optional, List, Any
from dataclasses import dataclass
import torch
import numpy as np
from app.core.accelerator import get_accelerator_info
from app.core.device import get_vram_gb
from app.core.hotword_resolver import format_hotword_prompt_context
from .engines import BaseASREngine, ASRRawResult, ASRSegmentResult, WordToken
from .qwenasr_rust import (
QwenASRRustRuntime,
is_qwenasr_rust_available,
resolve_qwenasr_model_path,
)
from .qwen3_vllm import Qwen3VLLMBackend, is_vllm_available
from ...core.exceptions import DefaultServerErrorException
from ...core.config import settings
from ...utils.text_processing import normalize_asr_text
logger = logging.getLogger(__name__)
def _resolve_tensor_parallel_size() -> int:
topology = (os.getenv("ASR_ACTIVE_TOPOLOGY") or os.getenv("ASR_DEPLOY_TOPOLOGY") or "isolated").strip().lower()
if topology != "sharded":
return 1
visible = (
os.getenv("ASR_VISIBLE_DEVICES")
or os.getenv("CUDA_VISIBLE_DEVICES")
or os.getenv("METAX_VISIBLE_DEVICES")
or os.getenv("MACA_VISIBLE_DEVICES")
or os.getenv("MX_VISIBLE_DEVICES")
or os.getenv("ILUVATAR_VISIBLE_DEVICES")
or os.getenv("IX_VISIBLE_DEVICES")
or os.getenv("MTHREADS_VISIBLE_DEVICES")
or os.getenv("MUSA_VISIBLE_DEVICES")
or ""
).strip()
if not visible or visible.lower() in {"all", "none", "void"}:
return 1
count = len([part for part in visible.split(",") if part.strip()])
return count if count > 1 else 1
def calculate_gpu_memory_utilization(model_path: str) -> float:
"""Calculate vLLM GPU memory utilization for the active model.
vLLM uses this ratio as an allocation budget, not just model weights.
Keep the observed requirement slightly above the bare minimum so KV cache
and profiling have enough room.
"""
# Check environment variable override first
env_override = os.getenv("QWEN_GPU_MEMORY_UTILIZATION")
if env_override:
try:
value = float(env_override)
if 0.0 < value <= 1.0:
logger.info(f"Using environment override: gpu_memory_utilization={value}")
return value
else:
logger.warning(f"Invalid QWEN_GPU_MEMORY_UTILIZATION={env_override}, must be 0.0-1.0")
except ValueError:
logger.warning(f"Invalid QWEN_GPU_MEMORY_UTILIZATION={env_override}, not a float")
model_memory_profiles = {
"0.6B": 8,
"1.7B": 12.0,
}
if "0.6B" in model_path:
model_size = "0.6B"
else:
model_size = "1.7B"
required_memory_gb = model_memory_profiles[model_size]
try:
accelerator = get_accelerator_info()
total_vram_gb = get_vram_gb()
if not accelerator.is_gpu or total_vram_gb <= 0:
logger.warning("Accelerator memory unavailable, using fallback gpu_memory_utilization=0.5")
return 0.5
utilization = max(required_memory_gb / total_vram_gb, 0.25)
utilization = min(utilization, 0.95)
logger.info(
"GPU memory calculation: vendor=%s, model=%s, requires=%.1fGB, total_vram=%.1fGB, utilization=%.2f",
accelerator.vendor,
model_size,
required_memory_gb,
total_vram_gb,
utilization,
)
if utilization >= 0.90:
logger.warning(
"VRAM may be insufficient: %.1fGB available, %.1fGB required. Consider using smaller model.",
total_vram_gb,
required_memory_gb,
)
return round(utilization, 2)
except Exception as e:
logger.error(f"Failed to detect VRAM: {e}, using fallback gpu_memory_utilization=0.5")
return 0.5
def _handle_asr_error(operation: str):
"""统一错误处理装饰器"""
def decorator(func):
def wrapper(*args, **kwargs):
try:
return func(*args, **kwargs)
except Exception as e:
logger.error(f"{operation} 失败: {e}")
raise DefaultServerErrorException(f"{operation} 失败: {e}")
return wrapper
return decorator
@dataclass
class Qwen3StreamingState:
internal_state: Any
chunk_size_sec: float = 1.2
unfixed_chunk_num: int = 2
unfixed_token_num: int = 5
max_new_tokens: int = 32
language: Optional[str] = None
chunk_count: int = 0
last_text: str = ""
last_language: str = ""
class Qwen3ASREngine(BaseASREngine):
model: Any
@property
def supports_realtime(self) -> bool:
return self._backend in {"vllm", "rust"}
def __init__(
self,
model_path: str = "Qwen/Qwen3-ASR-1.7B",
device: str = "auto",
forced_aligner_path: Optional[str] = None,
max_inference_batch_size: int = 32,
max_new_tokens: int = 1024,
max_model_len: Optional[int] = None,
**_kwargs,
):
"""Initialize Qwen3-ASR engine
CUDA -> official vLLM backend
CPU/macOS -> QwenASR Rust backend
"""
from app.core.device import detect_device
model_id = _kwargs.pop("model_id", None)
if model_id:
model_path = model_id
self._device = detect_device(device)
self._accelerator = get_accelerator_info()
self.model_id = model_path
self.model_path = model_path
self._backend = self._select_backend()
self._forced_aligner_path = forced_aligner_path
self._rust_num_threads = 0
self._rust_verbosity = 0
self._rust_batch_runtimes: list[QwenASRRustRuntime] = []
try:
if self._backend == "vllm":
self.model = self._load_vllm(
model_path, forced_aligner_path,
max_inference_batch_size, max_new_tokens, max_model_len,
)
elif self._backend == "rust":
self.model = self._load_rust_backend(model_path, forced_aligner_path)
self._warmup_forced_aligner()
logger.info("Qwen3-ASR model loaded successfully with backend=%s", self._backend)
except Exception as e:
logger.error(f"Failed to load Qwen3-ASR model: {e}")
raise DefaultServerErrorException(f"Failed to load Qwen3-ASR model: {e}")
def _select_backend(self) -> str:
if self._accelerator.is_gpu and self._device.startswith("cuda"):
if not is_vllm_available():
raise DefaultServerErrorException(
"Current Python environment is missing vLLM with Qwen3 forced aligner support. "
f"accelerator={self._accelerator.vendor}. "
"For NVIDIA run ./scripts/sync_gpu_env.sh; for MetaX run ./scripts/sync_metax_env.sh; "
"for Iluvatar or Moore Threads prefer the official vendor vLLM Docker image fused with this project."
)
return "vllm"
if self._device == "cpu" and is_qwenasr_rust_available():
return "rust"
raise DefaultServerErrorException(
f"Qwen3-ASR is not available on accelerator '{self._accelerator.vendor}' "
f"device '{self._device}'. Supported backends are vendor-compatible vLLM "
"GPU runtimes and CPU QwenASR Rust."
)
def _load_rust_backend(
self,
model_path: str,
forced_aligner_path: Optional[str],
) -> QwenASRRustRuntime:
logger.info("Loading Qwen3-ASR (QwenASR Rust): %s, device=%s", model_path, self._device)
num_threads = 0 if settings.QWEN_RUST_CPU_WORKERS <= 1 else 1
if settings.QWEN_RUST_CPU_WORKERS > 1:
logger.info(
"Using fixed QwenASR CPU thread count for multi-runtime mode: num_threads=%s workers=%s",
num_threads,
settings.QWEN_RUST_CPU_WORKERS,
)
self._rust_num_threads = num_threads
self._rust_verbosity = 0
return QwenASRRustRuntime(
model_path=model_path,
forced_aligner_path=forced_aligner_path,
num_threads=num_threads,
verbosity=0,
)
def _get_rust_batch_runtimes(self, worker_count: int) -> list[QwenASRRustRuntime]:
if worker_count <= 1:
return [self.model]
if not self._rust_batch_runtimes:
self._rust_batch_runtimes = [self.model]
while len(self._rust_batch_runtimes) < worker_count:
self._rust_batch_runtimes.append(
QwenASRRustRuntime(
model_path=self.model_path,
forced_aligner_path=self._forced_aligner_path,
num_threads=self._rust_num_threads,
verbosity=self._rust_verbosity,
)
)
return self._rust_batch_runtimes[:worker_count]
def _get_rust_stage_concurrency(self, configured: int, segment_count: int) -> int:
target = configured if configured > 0 else settings.QWEN_RUST_CPU_WORKERS
return max(1, min(target, segment_count))
def _get_rust_asr_concurrency(self, segment_count: int) -> int:
return self._get_rust_stage_concurrency(
settings.QWEN_RUST_ASR_CONCURRENCY,
segment_count,
)
def _get_rust_align_concurrency(self, segment_count: int) -> int:
return self._get_rust_stage_concurrency(
settings.QWEN_RUST_ALIGN_CONCURRENCY,
segment_count,
)
@staticmethod
def _build_hotword_prompt_context(hotwords: str) -> str:
return format_hotword_prompt_context(hotwords)
def _rust_transcribe_text_segment(
self,
runtime: QwenASRRustRuntime,
seg: Any,
hotwords: str,
enable_punctuation: bool,
enable_itn: bool,
sample_rate: int,
) -> str:
_ = (hotwords, enable_punctuation, sample_rate)
text = runtime.transcribe_file(seg.temp_file) or ""
return normalize_asr_text(text, enable_itn=enable_itn)
def _rust_align_word_tokens(
self,
runtime: QwenASRRustRuntime,
seg: Any,
text: str,
language: Optional[str] = None,
) -> list[WordToken]:
return [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in runtime.align_transcript(
audio_path=seg.temp_file,
text=text,
language=language,
)
]
def _run_rust_asr_stage(
self,
valid_segments: List[tuple[int, Any]],
hotwords: str,
enable_punctuation: bool,
enable_itn: bool,
sample_rate: int,
) -> dict[int, str]:
if not valid_segments:
return {}
worker_count = self._get_rust_asr_concurrency(len(valid_segments))
runtimes = self._get_rust_batch_runtimes(worker_count)
output: dict[int, str] = {}
for batch_start in range(0, len(valid_segments), worker_count):
chunk = valid_segments[batch_start:batch_start + worker_count]
chunk_runtimes = runtimes[:len(chunk)]
with ThreadPoolExecutor(max_workers=len(chunk)) as executor:
futures = [
executor.submit(
self._rust_transcribe_text_segment,
runtime,
seg,
hotwords,
enable_punctuation,
enable_itn,
sample_rate,
)
for runtime, (_idx, seg) in zip(chunk_runtimes, chunk)
]
for (idx, _seg), future in zip(chunk, futures):
output[idx] = future.result()
return output
def _run_rust_align_stage(
self,
valid_segments: List[tuple[int, Any]],
texts: dict[int, str],
language: Optional[str] = None,
) -> dict[int, list[WordToken]]:
if not valid_segments:
return {}
align_inputs = [(idx, seg, texts.get(idx, "")) for idx, seg in valid_segments if texts.get(idx, "").strip()]
worker_count = self._get_rust_align_concurrency(len(valid_segments))
runtimes = self._get_rust_batch_runtimes(worker_count)
output: dict[int, list[WordToken]] = {}
if not align_inputs:
return output
for batch_start in range(0, len(align_inputs), worker_count):
chunk = align_inputs[batch_start:batch_start + worker_count]
chunk_runtimes = runtimes[:len(chunk)]
with ThreadPoolExecutor(max_workers=len(chunk)) as executor:
futures = [
executor.submit(
self._rust_align_word_tokens,
runtime,
seg,
text,
language,
)
for runtime, (_idx, seg, text) in zip(chunk_runtimes, chunk)
]
for (idx, _seg, _text), future in zip(chunk, futures):
output[idx] = future.result()
return output
def _warmup_forced_aligner(self) -> None:
if not self._forced_aligner_path:
return
if not settings.ASR_ENABLE_WORD_TIMESTAMPS:
return
if self._backend == "vllm":
self.model.ensure_forced_aligner_loaded()
def _load_vllm(
self, model_path: str, forced_aligner_path: Optional[str],
max_inference_batch_size: int, max_new_tokens: int,
max_model_len: Optional[int],
) -> Qwen3VLLMBackend:
"""Load model via official vLLM backend (CUDA only)."""
resolved_model_path = str(resolve_qwenasr_model_path(model_path))
resolved_forced_aligner_path = None
if forced_aligner_path:
resolved_forced_aligner_path = str(resolve_qwenasr_model_path(forced_aligner_path))
gpu_memory_utilization = calculate_gpu_memory_utilization(model_path)
tensor_parallel_size = _resolve_tensor_parallel_size()
logger.info(
f"Loading Qwen3-ASR (official vLLM): {resolved_model_path}, "
f"device={self._device}, gpu_memory_utilization={gpu_memory_utilization}, "
f"enforce_eager={settings.QWEN_VLLM_ENFORCE_EAGER}, "
f"tensor_parallel_size={tensor_parallel_size}"
)
return Qwen3VLLMBackend(
model_path=resolved_model_path,
forced_aligner_path=resolved_forced_aligner_path,
gpu_memory_utilization=gpu_memory_utilization,
max_inference_batch_size=max_inference_batch_size,
max_new_tokens=max_new_tokens,
enforce_eager=settings.QWEN_VLLM_ENFORCE_EAGER,
max_model_len=max_model_len,
tensor_parallel_size=tensor_parallel_size,
)
@_handle_asr_error("转写")
def transcribe_file(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = True,
enable_itn: bool = True,
enable_vad: bool = False,
sample_rate: int = 16000,
) -> str:
if self._backend == "rust":
text = self.model.transcribe_file(audio_path)
return normalize_asr_text(text, enable_itn=enable_itn)
if self._backend == "vllm":
return self.model.transcribe_text(
audio_path,
context=self._build_hotword_prompt_context(hotwords),
enable_itn=enable_itn,
)
raise DefaultServerErrorException(f"Qwen3 backend={self._backend} does not support offline transcription")
@_handle_asr_error("VAD 转写")
def transcribe_file_with_vad(
self,
audio_path: str,
hotwords: str = "",
enable_punctuation: bool = True,
enable_itn: bool = True,
sample_rate: int = 16000,
**kwargs,
) -> ASRRawResult:
if self._backend == "rust":
text = self.transcribe_file(
audio_path=audio_path,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
)
if kwargs.get("word_timestamps", False):
word_tokens = [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in self.model.align_transcript(
audio_path=audio_path,
text=text,
language=kwargs.get("language"),
)
]
if word_tokens:
return ASRRawResult(
text=text,
segments=[
ASRSegmentResult(
text=text,
start_time=word_tokens[0].start_time,
end_time=word_tokens[-1].end_time,
word_tokens=word_tokens,
)
],
)
return ASRRawResult(
text=text,
segments=[ASRSegmentResult(text=text, start_time=0.0, end_time=0.0)] if text else [],
)
if self._backend == "vllm":
return self.model.transcribe_raw(
audio_path=audio_path,
context=self._build_hotword_prompt_context(hotwords),
language=kwargs.get("language"),
word_timestamps=kwargs.get("word_timestamps", False),
enable_itn=enable_itn,
)
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support VAD transcription"
)
@_handle_asr_error("批量推理")
def _transcribe_batch(
self,
segments: List[Any],
hotwords: str = "",
enable_punctuation: bool = False,
enable_itn: bool = False,
sample_rate: int = 16000,
word_timestamps: bool = False,
) -> List[ASRSegmentResult]:
output = [ASRSegmentResult(text="", start_time=0.0, end_time=0.0) for _ in segments]
valid: List[tuple[int, Any]] = []
for idx, seg in enumerate(segments):
temp_file = getattr(seg, "temp_file", None)
if temp_file and os.path.exists(temp_file):
valid.append((idx, seg))
else:
logger.warning(f"Qwen3 批处理片段无效或文件不存在: segment={idx + 1}, file={temp_file}")
if not valid:
return output
if self._backend == "rust":
texts = self._run_rust_asr_stage(
valid_segments=valid,
hotwords=hotwords,
enable_punctuation=enable_punctuation,
enable_itn=enable_itn,
sample_rate=sample_rate,
)
word_tokens_by_idx: dict[int, list[WordToken]] = {}
if word_timestamps:
word_tokens_by_idx = self._run_rust_align_stage(
valid_segments=valid,
texts=texts,
)
for idx, seg in valid:
text = texts.get(idx, "")
output[idx] = ASRSegmentResult(
text=text,
start_time=seg.start_sec,
end_time=seg.end_sec,
speaker_id=getattr(seg, "speaker_id", None),
word_tokens=word_tokens_by_idx.get(idx) or None,
)
return output
if self._backend == "vllm":
vllm_results = self.model.transcribe_batch(
[seg.temp_file for _, seg in valid],
context=self._build_hotword_prompt_context(hotwords),
word_timestamps=word_timestamps,
enable_itn=enable_itn,
)
for (idx, seg), result in zip(valid, vllm_results):
output[idx] = ASRSegmentResult(
text=result.text,
start_time=round(seg.start_sec, 2),
end_time=round(seg.end_sec, 2),
speaker_id=getattr(seg, "speaker_id", None),
word_tokens=result.word_tokens if word_timestamps else None,
)
return output
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support batch transcription"
)
@_handle_asr_error("初始化流式状态")
def init_streaming_state(self, context: str = "", language: Optional[str] = None, **kwargs) -> Qwen3StreamingState:
if self._backend not in {"vllm", "rust"}:
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support realtime streaming"
)
if self._backend == "rust":
if context:
logger.debug("QwenASR Rust backend ignores streaming context hints")
chunk_size_sec = float(kwargs.get("chunk_size_sec", 1.2))
unfixed_chunk_num = int(kwargs.get("unfixed_chunk_num", 2))
unfixed_token_num = int(kwargs.get("unfixed_token_num", 5))
max_new_tokens = int(kwargs.get("max_new_tokens", 32))
stream_handle = self.model.create_stream(
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=unfixed_token_num,
max_new_tokens=max_new_tokens,
language=language,
)
return Qwen3StreamingState(
internal_state=stream_handle,
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
unfixed_token_num=unfixed_token_num,
max_new_tokens=max_new_tokens,
language=language,
chunk_count=0,
last_text="",
last_language=language or "",
)
if self._backend == "vllm":
streaming_state = self.model.init_streaming_state(context=context, language=language, **kwargs)
return Qwen3StreamingState(
internal_state=streaming_state,
chunk_size_sec=float(kwargs.get("chunk_size_sec", 1.2)),
unfixed_chunk_num=int(kwargs.get("unfixed_chunk_num", 2)),
unfixed_token_num=int(kwargs.get("unfixed_token_num", 5)),
max_new_tokens=int(kwargs.get("max_new_tokens", 32)),
language=language,
chunk_count=int(getattr(streaming_state, "chunk_id", 0)),
last_text=str(getattr(streaming_state, "text", "") or ""),
last_language=str(getattr(streaming_state, "language", "") or ""),
)
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support realtime streaming"
)
@_handle_asr_error("流式识别")
def streaming_transcribe(self, pcm16k: np.ndarray, state: Qwen3StreamingState) -> Qwen3StreamingState:
if self._backend not in {"vllm", "rust"}:
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support realtime streaming"
)
pcm = pcm16k.astype(np.float32) / (32768.0 if pcm16k.dtype == np.int16 else 1.0)
if self._backend == "rust":
text = self.model.push_stream(
stream=state.internal_state,
samples=pcm,
chunk_size_sec=state.chunk_size_sec,
unfixed_chunk_num=state.unfixed_chunk_num,
rollback_tokens=state.unfixed_token_num,
max_new_tokens=state.max_new_tokens,
language=state.language,
)
state.chunk_count += 1
state.last_text = text
state.last_language = state.language or ""
return state
streaming_state = self.model.feed_stream(pcm, state.internal_state)
state.internal_state = streaming_state
state.chunk_count = int(getattr(streaming_state, "chunk_id", state.chunk_count))
state.last_text = str(getattr(streaming_state, "text", "") or "")
state.last_language = str(getattr(streaming_state, "language", "") or "")
return state
@_handle_asr_error("结束流式识别")
def finish_streaming_transcribe(self, state: Qwen3StreamingState) -> Qwen3StreamingState:
if self._backend not in {"vllm", "rust"}:
raise DefaultServerErrorException(
f"Qwen3 backend={self._backend} does not support realtime streaming"
)
if self._backend == "rust":
text = self.model.finish_stream(
stream=state.internal_state,
chunk_size_sec=state.chunk_size_sec,
unfixed_chunk_num=state.unfixed_chunk_num,
rollback_tokens=state.unfixed_token_num,
max_new_tokens=state.max_new_tokens,
language=state.language,
)
state.last_text = text
state.last_language = state.language or ""
return state
streaming_state = self.model.finish_stream(state.internal_state)
state.internal_state = streaming_state
state.chunk_count = int(getattr(streaming_state, "chunk_id", state.chunk_count))
state.last_text = str(getattr(streaming_state, "text", "") or "")
state.last_language = str(getattr(streaming_state, "language", "") or "")
return state
def is_model_loaded(self) -> bool:
return self.model is not None
@property
def backend(self) -> str:
return self._backend
@property
def device(self) -> str:
return self._device
def _register_qwen3_engine(register_func, _declared_entry_cls):
from app.core.config import settings
def _create(config):
extra = {k: v for k, v in config.extra_kwargs.items() if v is not None}
model_id = config.models.get("offline")
return Qwen3ASREngine(model_path=model_id, device=settings.DEVICE, **extra)
register_func("qwen3", _create)

View File

@ -0,0 +1,527 @@
# -*- coding: utf-8 -*-
"""Official vLLM adapter for CUDA Qwen3-ASR."""
from __future__ import annotations
import importlib
import importlib.util
import logging
import os
import re
import threading
from dataclasses import dataclass, field
from typing import Any, Optional
import librosa
import numpy as np
from app.core.hotword_resolver import strip_hotword_prompt_leakage
from app.utils.text_processing import normalize_asr_text
from .engines import ASRRawResult, ASRSegmentResult, WordToken
logger = logging.getLogger(__name__)
_DEFAULT_SAMPLE_RATE = 16000
_LANGUAGE_ALIASES = {
"zh": "Chinese",
"zh-cn": "Chinese",
"zh-hans": "Chinese",
"zh-hant": "Chinese",
"cn": "Chinese",
"en": "English",
"en-us": "English",
"en-gb": "English",
"ja": "Japanese",
"jp": "Japanese",
"ko": "Korean",
"yue": "Cantonese",
"fr": "French",
"de": "German",
"es": "Spanish",
"ru": "Russian",
}
_PROMPT_LEAK_PATTERNS = (
re.compile(r"^\s*Transcribe the speech accurately\.\s*", re.IGNORECASE),
re.compile(r"^\s*Transcribe the speech in [A-Za-z\s-]+\.\s*", re.IGNORECASE),
)
def is_vllm_available() -> bool:
"""Return True when the official vLLM runtime is installed."""
return importlib.util.find_spec("vllm") is not None
def _normalize_language_name(language: Optional[str]) -> Optional[str]:
if not language:
return None
normalized = language.strip()
if not normalized:
return None
alias = _LANGUAGE_ALIASES.get(normalized.lower())
if alias:
return alias
if " " in normalized:
return " ".join(part.capitalize() for part in normalized.split())
return normalized.capitalize()
def _load_audio(audio_path: str) -> np.ndarray:
audio, _sample_rate = librosa.load(audio_path, sr=_DEFAULT_SAMPLE_RATE, mono=True)
return audio.astype(np.float32)
def _build_chat_prompt(context: str = "", language: Optional[str] = None) -> str:
instructions: list[str] = []
if language:
instructions.append(f"Transcribe the speech in {language}.")
else:
instructions.append("Transcribe the speech accurately.")
if context.strip():
instructions.append(context.strip())
system_text = " ".join(instructions).strip()
return (
f"<|im_start|>system\n{system_text}<|im_end|>\n"
"<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|><|im_end|>\n"
"<|im_start|>assistant\n"
)
def _build_alignment_prompt(tokens: list[str]) -> str:
body = "<timestamp><timestamp>".join(tokens) + "<timestamp><timestamp>"
return f"<|audio_start|><|audio_pad|><|audio_end|>{body}"
def _strip_prompt_leakage(text: str) -> str:
cleaned = text or ""
changed = True
while changed and cleaned:
changed = False
for pattern in _PROMPT_LEAK_PATTERNS:
updated, count = pattern.subn("", cleaned, count=1)
if count:
cleaned = updated
changed = True
return strip_hotword_prompt_leakage(cleaned)
def _sanitize_detected_language(detected: str, fallback: Optional[str]) -> str:
candidate = (detected or "").strip()
if not candidate:
return fallback or ""
candidate = re.sub(r"^[^\w]+", "", candidate)
match = re.search(r"language\s+([A-Za-z][A-Za-z\s-]*)$", candidate, re.IGNORECASE)
if match:
candidate = match.group(1).strip()
normalized = _normalize_language_name(candidate)
return normalized or (fallback or "")
def _parse_asr_output(raw_text: str, language: Optional[str]) -> tuple[str, str]:
text = (raw_text or "").strip()
if "<asr_text>" in text:
left, right = text.split("<asr_text>", 1)
detected = _sanitize_detected_language(left.strip(), language)
return detected, _strip_prompt_leakage(right.strip())
return (language or ""), _strip_prompt_leakage(text)
def _split_alignment_units(text: str) -> list[str]:
if not text:
return []
# Mixed Chinese/English transcripts should not fall back to whitespace-only
# tokenization, otherwise a long CJK sentence with a single embedded English
# word can collapse into one giant alignment unit.
token_pattern = re.compile(
r"[\u4e00-\u9fff]" # CJK ideographs, align per character
r"|[A-Za-z0-9]+(?:['._+-][A-Za-z0-9]+)*" # Latin / alnum words
r"|[^\w\s]", # punctuation and symbols
re.UNICODE,
)
return token_pattern.findall(text)
def _resolve_forced_aligner_gpu_memory_utilization(primary_utilization: float) -> float:
override = (os.getenv("QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION") or "").strip()
if override:
try:
value = float(override)
if 0.0 < value <= 1.0:
return value
except ValueError:
logger.warning(
"Invalid QWEN_FORCE_ALIGNER_GPU_MEMORY_UTILIZATION=%s, ignoring override",
override,
)
return primary_utilization
@dataclass
class _GeneratedTranscript:
text: str
language: str
@dataclass
class VLLMRealtimeState:
prompt_raw: str
language: str
chunk_size_sec: float
unfixed_chunk_num: int
unfixed_token_num: int
max_new_tokens: int
chunk_id: int = 0
text: str = ""
raw_decoded: str = ""
audio_buffer: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32))
audio_accum: np.ndarray = field(default_factory=lambda: np.array([], dtype=np.float32))
class Qwen3VLLMBackend:
"""Thin adapter over official vLLM APIs for Qwen3-ASR."""
def __init__(
self,
model_path: str,
forced_aligner_path: Optional[str],
gpu_memory_utilization: float,
max_inference_batch_size: int,
max_new_tokens: int,
enforce_eager: bool = True,
max_model_len: Optional[int] = None,
tensor_parallel_size: int = 1,
) -> None:
try:
vllm_module = importlib.import_module("vllm")
transformers_module = importlib.import_module("transformers")
except ImportError as exc:
raise RuntimeError(
"CUDA Qwen3-ASR now requires official vLLM with Qwen3 forced aligner support. "
"Install it with: pip install 'vllm[audio]==0.19.0'"
) from exc
self._llm_cls = getattr(vllm_module, "LLM")
self._sampling_params_cls = getattr(vllm_module, "SamplingParams")
self._tokenizer = getattr(transformers_module, "AutoTokenizer").from_pretrained(
model_path,
trust_remote_code=True,
)
llm_kwargs: dict[str, Any] = {
"model": model_path,
"gpu_memory_utilization": gpu_memory_utilization,
"enforce_eager": enforce_eager,
"trust_remote_code": True,
}
if max_model_len is not None:
llm_kwargs["max_model_len"] = max_model_len
if tensor_parallel_size > 1:
llm_kwargs["tensor_parallel_size"] = tensor_parallel_size
self._llm = self._llm_cls(**llm_kwargs)
self._sampling_params = self._sampling_params_cls(
temperature=0.01,
max_tokens=max_new_tokens,
)
self._max_inference_batch_size = max_inference_batch_size
self._gpu_memory_utilization = gpu_memory_utilization
self._enforce_eager = enforce_eager
self._tensor_parallel_size = tensor_parallel_size
self._forced_aligner_path = forced_aligner_path
self._forced_aligner: Any | None = None
self._timestamp_token_id: int | None = None
self._timestamp_segment_time: float | None = None
# Share one backend instance across multiple offline tasks, but serialize
# direct vLLM engine calls to avoid cross-request state corruption/hangs.
self._engine_lock = threading.RLock()
def _get_forced_aligner_gpu_memory_utilization(self) -> float:
configured = _resolve_forced_aligner_gpu_memory_utilization(self._gpu_memory_utilization)
logger.info(
"Resolved forced aligner gpu_memory_utilization=%s (primary=%s)",
configured,
self._gpu_memory_utilization,
)
return configured
def _get_forced_aligner(self) -> Any:
if not self._forced_aligner_path:
raise RuntimeError("word_timestamps requires a configured forced aligner model")
if self._forced_aligner is None:
forced_aligner_gpu_memory_utilization = self._get_forced_aligner_gpu_memory_utilization()
logger.info(
"Loading Qwen3 forced aligner via official vLLM: %s (gpu_memory_utilization=%s)",
self._forced_aligner_path,
forced_aligner_gpu_memory_utilization,
)
self._forced_aligner = self._llm_cls(
model=self._forced_aligner_path,
runner="pooling",
enforce_eager=self._enforce_eager,
gpu_memory_utilization=forced_aligner_gpu_memory_utilization,
tensor_parallel_size=self._tensor_parallel_size,
trust_remote_code=True,
hf_overrides={
"architectures": ["Qwen3ASRForcedAlignerForTokenClassification"],
},
)
llm_engine = getattr(self._forced_aligner, "llm_engine", None)
if llm_engine is None:
raise RuntimeError("Forced aligner did not expose a vLLM engine instance")
config = llm_engine.vllm_config.model_config.hf_config
self._timestamp_token_id = int(config.timestamp_token_id)
self._timestamp_segment_time = float(config.timestamp_segment_time)
return self._forced_aligner
def ensure_forced_aligner_loaded(self) -> None:
if self._forced_aligner_path:
self._get_forced_aligner()
def _run_generate(
self,
audio_items: list[tuple[np.ndarray, str, Optional[str]]],
) -> list[_GeneratedTranscript]:
prompts: list[dict[str, Any]] = []
for audio, context, language in audio_items:
prompts.append(
{
"prompt": _build_chat_prompt(context=context, language=_normalize_language_name(language)),
"multi_modal_data": {"audio": [audio]},
}
)
with self._engine_lock:
outputs = self._llm.generate(
prompts,
sampling_params=self._sampling_params,
use_tqdm=False,
)
transcripts: list[_GeneratedTranscript] = []
for output, (_audio, _context, language) in zip(outputs, audio_items):
raw_text = str(output.outputs[0].text if output.outputs else "")
parsed_language, parsed_text = _parse_asr_output(raw_text, _normalize_language_name(language))
transcripts.append(_GeneratedTranscript(text=parsed_text, language=parsed_language))
return transcripts
def transcribe_text(
self,
audio_path: str,
context: str = "",
language: Optional[str] = None,
enable_itn: bool = False,
) -> str:
transcript = self._run_generate([(_load_audio(audio_path), context, language)])[0]
return normalize_asr_text(transcript.text, enable_itn=enable_itn)
def transcribe_raw(
self,
audio_path: str,
context: str = "",
language: Optional[str] = None,
word_timestamps: bool = False,
enable_itn: bool = False,
) -> ASRRawResult:
audio = _load_audio(audio_path)
transcript = self._run_generate([(audio, context, language)])[0]
text = normalize_asr_text(transcript.text, enable_itn=enable_itn)
if not word_timestamps:
return ASRRawResult(
text=text,
segments=[ASRSegmentResult(text=text, start_time=0.0, end_time=0.0)] if text else [],
)
aligned = self.align_transcript(audio_path=audio_path, text=text, language=language, audio=audio)
word_tokens = [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in aligned
]
if not word_tokens:
return ASRRawResult(
text=text,
segments=[ASRSegmentResult(text=text, start_time=0.0, end_time=0.0)] if text else [],
)
return ASRRawResult(
text=text,
segments=[
ASRSegmentResult(
text=text,
start_time=word_tokens[0].start_time,
end_time=word_tokens[-1].end_time,
word_tokens=word_tokens,
)
],
)
def transcribe_batch(
self,
audio_paths: list[str],
context: str = "",
language: Optional[str] = None,
word_timestamps: bool = False,
enable_itn: bool = False,
) -> list[ASRSegmentResult]:
audios = [_load_audio(path) for path in audio_paths]
results: list[ASRSegmentResult] = []
for start in range(0, len(audios), self._max_inference_batch_size):
chunk = audios[start:start + self._max_inference_batch_size]
transcripts = self._run_generate([(audio, context, language) for audio in chunk])
for audio_path, audio, transcript in zip(audio_paths[start:start + len(chunk)], chunk, transcripts):
text = normalize_asr_text(transcript.text, enable_itn=enable_itn)
if not word_timestamps:
results.append(ASRSegmentResult(text=text, start_time=0.0, end_time=0.0))
continue
aligned = self.align_transcript(
audio_path=audio_path,
text=text,
language=language,
audio=audio,
)
word_tokens = [
WordToken(
text=str(item["text"]),
start_time=round(float(item["start_ms"]) / 1000.0, 3),
end_time=round(float(item["end_ms"]) / 1000.0, 3),
)
for item in aligned
]
results.append(
ASRSegmentResult(
text=text,
start_time=word_tokens[0].start_time if word_tokens else 0.0,
end_time=word_tokens[-1].end_time if word_tokens else 0.0,
word_tokens=word_tokens or None,
)
)
return results
def align_transcript(
self,
audio_path: str,
text: str,
language: Optional[str] = None,
audio: Optional[np.ndarray] = None,
) -> list[dict[str, float | str]]:
tokens = _split_alignment_units(text)
if not tokens:
return []
aligner = self._get_forced_aligner()
prompt = _build_alignment_prompt(tokens)
audio_array = audio if audio is not None else _load_audio(audio_path)
with self._engine_lock:
outputs = aligner.encode(
[{"prompt": prompt, "multi_modal_data": {"audio": audio_array}}],
pooling_task="token_classify",
)
output = outputs[0]
logits = output.outputs.data
predictions = logits.argmax(dim=-1) if hasattr(logits, "argmax") else np.argmax(logits, axis=-1)
ts_predictions = [
float(pred.item() if hasattr(pred, "item") else pred) * float(self._timestamp_segment_time or 0.0)
for tid, pred in zip(output.prompt_token_ids, predictions)
if int(tid) == int(self._timestamp_token_id or -1)
]
expected_timestamps = len(tokens) * 2
if len(ts_predictions) < expected_timestamps:
raise RuntimeError(
"Forced aligner returned fewer timestamp predictions than expected: "
f"expected={expected_timestamps}, got={len(ts_predictions)}, tokens={len(tokens)}"
)
aligned: list[dict[str, float | str]] = []
for index, token in enumerate(tokens):
start_ms = ts_predictions[index * 2]
end_ms = ts_predictions[index * 2 + 1]
if end_ms < start_ms:
logger.warning(
"Forced aligner produced reversed timestamps for token=%r: start_ms=%s end_ms=%s",
token,
start_ms,
end_ms,
)
start_ms, end_ms = end_ms, start_ms
aligned.append({"text": token, "start_ms": start_ms, "end_ms": end_ms})
return aligned
def init_streaming_state(
self,
*,
context: str = "",
language: Optional[str] = None,
chunk_size_sec: float = 1.2,
unfixed_chunk_num: int = 2,
unfixed_token_num: int = 5,
max_new_tokens: int = 32,
) -> VLLMRealtimeState:
normalized_language = _normalize_language_name(language) or ""
return VLLMRealtimeState(
prompt_raw=_build_chat_prompt(context=context, language=normalized_language or None),
language=normalized_language,
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
unfixed_token_num=unfixed_token_num,
max_new_tokens=max_new_tokens,
audio_buffer=np.array([], dtype=np.float32),
audio_accum=np.array([], dtype=np.float32),
)
def _decode_stream(self, state: VLLMRealtimeState) -> VLLMRealtimeState:
prefix = ""
if state.chunk_id >= state.unfixed_chunk_num and state.raw_decoded:
token_ids = self._tokenizer.encode(state.raw_decoded, add_special_tokens=False)
rollback = token_ids[-state.unfixed_token_num:] if state.unfixed_token_num > 0 else []
if rollback:
prefix = self._tokenizer.decode(rollback, skip_special_tokens=False).replace("\ufffd", "")
with self._engine_lock:
output = self._llm.generate(
[
{
"prompt": state.prompt_raw + prefix,
"multi_modal_data": {"audio": [state.audio_accum]},
}
],
sampling_params=self._sampling_params_cls(
temperature=0.01,
max_tokens=state.max_new_tokens,
),
use_tqdm=False,
)[0]
generated = str(output.outputs[0].text if output.outputs else "")
parsed_language, parsed_text = _parse_asr_output(prefix + generated, state.language or None)
state.raw_decoded = prefix + generated
state.text = parsed_text
state.language = parsed_language or state.language
state.chunk_id += 1
return state
def feed_stream(self, pcm: np.ndarray, state: VLLMRealtimeState) -> VLLMRealtimeState:
state.audio_buffer = np.concatenate([state.audio_buffer, pcm.astype(np.float32)])
segment_size = int(max(state.chunk_size_sec, 0.1) * _DEFAULT_SAMPLE_RATE)
while len(state.audio_buffer) >= segment_size:
segment = state.audio_buffer[:segment_size].copy()
state.audio_buffer = state.audio_buffer[segment_size:]
state.audio_accum = np.concatenate([state.audio_accum, segment])
state = self._decode_stream(state)
return state
def finish_stream(self, state: VLLMRealtimeState) -> VLLMRealtimeState:
if len(state.audio_buffer) > 0:
state.audio_accum = np.concatenate([state.audio_accum, state.audio_buffer])
state.audio_buffer = np.array([], dtype=np.float32)
state = self._decode_stream(state)
elif state.chunk_id == 0 and len(state.audio_accum) > 0:
state = self._decode_stream(state)
return state

View File

@ -0,0 +1,564 @@
# -*- coding: utf-8 -*-
"""QwenASR Rust FFI wrapper for CPU inference."""
from __future__ import annotations
import ctypes
import json
import logging
import os
import platform
import re
import sys
from pathlib import Path
from typing import Optional
import numpy as np
from app.core.config import settings
logger = logging.getLogger(__name__)
_SHARED_LIBRARY: Optional[ctypes.CDLL] = None
_LANGUAGE_MAP = {
"": "",
"auto": "",
"zh": "Chinese",
"zh-cn": "Chinese",
"yue": "Chinese",
"en": "English",
"ja": "Japanese",
"ko": "Korean",
"de": "German",
"es": "Spanish",
"fr": "French",
"it": "Italian",
"pt": "Portuguese",
"ru": "Russian",
"ar": "Arabic",
"th": "Thai",
"vi": "Vietnamese",
"id": "Indonesian",
}
def _shared_library_filename() -> str:
if sys.platform == "darwin":
return "libqwen_asr.dylib"
if sys.platform == "win32":
return "qwen_asr.dll"
return "libqwen_asr.so"
def _repo_root() -> Path:
return Path(__file__).resolve().parents[3]
def _candidate_library_paths() -> list[Path]:
filename = _shared_library_filename()
candidates: list[Path] = []
env_path = (os.getenv("QWENASR_LIBRARY_PATH") or "").strip()
if env_path:
candidate = Path(env_path).expanduser()
if candidate.is_dir():
candidates.append(candidate / filename)
else:
candidates.append(candidate)
repo_root = _repo_root()
candidates.extend(
[
repo_root / "vendor" / "qwenasr" / "target" / "release" / filename,
repo_root / "vendor" / "qwenasr" / "target" / "debug" / filename,
Path("/opt/qwenasr/lib") / filename,
Path("/usr/local/lib") / filename,
]
)
return candidates
def resolve_qwenasr_library_path() -> Optional[Path]:
for candidate in _candidate_library_paths():
if candidate.exists():
return candidate.resolve()
return None
def is_qwenasr_rust_available() -> bool:
return resolve_qwenasr_library_path() is not None
def validate_qwenasr_cpu_features() -> None:
if platform.machine().lower() not in {"amd64", "x86_64"}:
return
flags = _read_linux_cpu_flags()
if not flags:
return
missing = [flag for flag in ("avx2", "fma") if flag not in flags]
if missing:
raise RuntimeError(
"QwenASR Rust backend requires x86_64 CPU features: avx2, fma. "
f"Missing: {', '.join(missing)}. Use a newer CPU host or rebuild the "
"Rust backend with scalar x86 kernels."
)
def pick_cpu_qwen_model(all_available_models: list[str]) -> Optional[str]:
for model_id in ["qwen3-asr-0.6b", "qwen3-asr-1.7b"]:
if model_id in all_available_models:
return model_id
return None
def _read_linux_cpu_flags() -> set[str]:
cpuinfo = Path("/proc/cpuinfo")
if not cpuinfo.exists():
return set()
flags: set[str] = set()
for line in cpuinfo.read_text(encoding="utf-8", errors="ignore").splitlines():
key, _, value = line.partition(":")
if key.strip().lower() in {"flags", "features"}:
flags.update(value.strip().lower().split())
if flags:
break
return flags
def _resolve_modelscope_dir(model_ref: str, cache_root: Path) -> Optional[Path]:
if "/" not in model_ref:
return None
base_dir = cache_root / model_ref
if base_dir.exists() and base_dir.is_dir():
return base_dir.resolve()
return None
def _append_unique_path(paths: list[Path], path: Path) -> None:
if path not in paths:
paths.append(path)
def resolve_qwenasr_model_path(model_ref_or_path: str) -> Path:
raw_path = Path(model_ref_or_path).expanduser()
if raw_path.exists():
return raw_path.resolve()
modelscope_cache_roots: list[Path] = []
ms_cache = (os.getenv("MODELSCOPE_CACHE") or "").strip()
if ms_cache:
cache_root = Path(ms_cache).expanduser()
_append_unique_path(modelscope_cache_roots, cache_root)
# 兼容旧目录结构,允许外部仍传入 models/modelscope
legacy_cache_root = cache_root / "hub" / "models"
if legacy_cache_root.exists():
_append_unique_path(modelscope_cache_roots, legacy_cache_root)
default_ms_cache_root = Path(settings.MODELSCOPE_PATH).expanduser()
_append_unique_path(modelscope_cache_roots, default_ms_cache_root)
for cache_root in modelscope_cache_roots:
ms_dir = _resolve_modelscope_dir(model_ref_or_path, cache_root)
if ms_dir is not None:
return ms_dir
raise FileNotFoundError(
f"QwenASR model path not found for '{model_ref_or_path}'. "
f"Checked direct path and ModelScope caches at: "
f"{', '.join(str(path) for path in modelscope_cache_roots)}."
)
def _bind_ffi_signatures(lib: ctypes.CDLL) -> None:
lib.qwen_asr_load_model.argtypes = [ctypes.c_char_p, ctypes.c_int, ctypes.c_int]
lib.qwen_asr_load_model.restype = ctypes.c_void_p
lib.qwen_asr_transcribe_file.argtypes = [ctypes.c_void_p, ctypes.c_char_p]
lib.qwen_asr_transcribe_file.restype = ctypes.c_void_p
lib.qwen_asr_force_align_file.argtypes = [
ctypes.c_void_p,
ctypes.c_char_p,
ctypes.c_char_p,
ctypes.c_char_p,
]
lib.qwen_asr_force_align_file.restype = ctypes.c_void_p
lib.qwen_asr_set_language.argtypes = [ctypes.c_void_p, ctypes.c_char_p]
lib.qwen_asr_set_language.restype = ctypes.c_int
lib.qwen_asr_free_string.argtypes = [ctypes.c_void_p]
lib.qwen_asr_free_string.restype = None
lib.qwen_asr_free.argtypes = [ctypes.c_void_p]
lib.qwen_asr_free.restype = None
lib.qwen_asr_stream_new.argtypes = []
lib.qwen_asr_stream_new.restype = ctypes.c_void_p
lib.qwen_asr_stream_free.argtypes = [ctypes.c_void_p]
lib.qwen_asr_stream_free.restype = None
lib.qwen_asr_stream_push.argtypes = [
ctypes.c_void_p,
ctypes.c_void_p,
ctypes.POINTER(ctypes.c_float),
ctypes.c_int,
ctypes.c_int,
]
lib.qwen_asr_stream_push.restype = ctypes.c_void_p
lib.qwen_asr_stream_get_result.argtypes = [ctypes.c_void_p]
lib.qwen_asr_stream_get_result.restype = ctypes.c_void_p
lib.qwen_asr_stream_set_chunk_sec.argtypes = [ctypes.c_void_p, ctypes.c_float]
lib.qwen_asr_stream_set_chunk_sec.restype = None
lib.qwen_asr_stream_set_rollback.argtypes = [ctypes.c_void_p, ctypes.c_int]
lib.qwen_asr_stream_set_rollback.restype = None
lib.qwen_asr_stream_set_unfixed_chunks.argtypes = [ctypes.c_void_p, ctypes.c_int]
lib.qwen_asr_stream_set_unfixed_chunks.restype = None
lib.qwen_asr_stream_set_max_new_tokens.argtypes = [ctypes.c_void_p, ctypes.c_int]
lib.qwen_asr_stream_set_max_new_tokens.restype = None
lib.qwen_asr_stream_set_past_text.argtypes = [ctypes.c_void_p, ctypes.c_int]
lib.qwen_asr_stream_set_past_text.restype = None
def load_qwenasr_library() -> ctypes.CDLL:
global _SHARED_LIBRARY
if _SHARED_LIBRARY is not None:
return _SHARED_LIBRARY
library_path = resolve_qwenasr_library_path()
if library_path is None:
searched = ", ".join(str(path) for path in _candidate_library_paths())
raise FileNotFoundError(
"QwenASR shared library not found. "
f"Checked: {searched}"
)
logger.info("Loading QwenASR Rust library from %s", library_path)
library = ctypes.CDLL(str(library_path))
_bind_ffi_signatures(library)
_SHARED_LIBRARY = library
return library
def normalize_qwen_language(language: Optional[str]) -> str:
if language is None:
return ""
return _LANGUAGE_MAP.get(language.strip().lower(), language.strip())
def guess_alignment_language(text: str, language: Optional[str] = None) -> str:
normalized = normalize_qwen_language(language)
if normalized:
return normalized
if re.search(r"[\u4e00-\u9fff]", text):
return "Chinese"
if re.search(r"[\u3040-\u30ff]", text):
return "Japanese"
if re.search(r"[\uac00-\ud7af]", text):
return "Korean"
return "English"
def _decode_and_free_string(lib: ctypes.CDLL, raw_ptr: ctypes.c_void_p) -> Optional[str]:
if not raw_ptr:
return None
try:
value = ctypes.cast(raw_ptr, ctypes.c_char_p).value
if value is None:
return None
return value.decode("utf-8")
finally:
lib.qwen_asr_free_string(raw_ptr)
class QwenASRRustStreamHandle:
"""Owns a Rust streaming state pointer."""
def __init__(self, lib: ctypes.CDLL, handle: ctypes.c_void_p):
self._lib = lib
self.handle = handle
self.accumulated_text = ""
def close(self) -> None:
if self.handle:
self._lib.qwen_asr_stream_free(self.handle)
self.handle = ctypes.c_void_p()
def __del__(self) -> None:
try:
self.close()
except Exception:
pass
class QwenASRRustBackend:
"""Thin Python wrapper around the QwenASR C API."""
def __init__(self, model_path: str, num_threads: int = 0, verbosity: int = 0):
validate_qwenasr_cpu_features()
self._lib = load_qwenasr_library()
self.model_dir = resolve_qwenasr_model_path(model_path)
self._engine = self._lib.qwen_asr_load_model(
str(self.model_dir).encode("utf-8"),
num_threads,
verbosity,
)
if not self._engine:
raise RuntimeError(f"Failed to load QwenASR model from '{self.model_dir}'")
def close(self) -> None:
if self._engine:
self._lib.qwen_asr_free(self._engine)
self._engine = ctypes.c_void_p()
def __del__(self) -> None:
try:
self.close()
except Exception:
pass
def _set_language(self, language: Optional[str]) -> None:
normalized = normalize_qwen_language(language)
status = self._lib.qwen_asr_set_language(
self._engine,
normalized.encode("utf-8"),
)
if status != 0 and normalized:
logger.warning("QwenASR rejected language hint: %s", normalized)
def _configure_stream(
self,
*,
chunk_size_sec: float,
unfixed_chunk_num: int,
rollback_tokens: int,
max_new_tokens: int,
past_text: bool,
) -> None:
self._lib.qwen_asr_stream_set_chunk_sec(self._engine, ctypes.c_float(chunk_size_sec))
self._lib.qwen_asr_stream_set_unfixed_chunks(self._engine, int(unfixed_chunk_num))
self._lib.qwen_asr_stream_set_rollback(self._engine, int(rollback_tokens))
self._lib.qwen_asr_stream_set_max_new_tokens(self._engine, int(max_new_tokens))
self._lib.qwen_asr_stream_set_past_text(self._engine, 1 if past_text else 0)
def transcribe_file(self, audio_path: str, language: Optional[str] = None) -> str:
self._set_language(language)
raw_ptr = self._lib.qwen_asr_transcribe_file(
self._engine,
audio_path.encode("utf-8"),
)
text = _decode_and_free_string(self._lib, raw_ptr)
if text is None:
raise RuntimeError(f"QwenASR failed to transcribe '{audio_path}'")
return text
def force_align_file(
self,
audio_path: str,
text: str,
language: Optional[str] = None,
) -> list[dict[str, float | str]]:
normalized_language = normalize_qwen_language(language) or "English"
raw_ptr = self._lib.qwen_asr_force_align_file(
self._engine,
audio_path.encode("utf-8"),
text.encode("utf-8"),
normalized_language.encode("utf-8"),
)
payload = _decode_and_free_string(self._lib, raw_ptr)
if payload is None:
raise RuntimeError(f"QwenASR failed to force align '{audio_path}'")
items = json.loads(payload)
if not isinstance(items, list):
raise RuntimeError("QwenASR force alignment returned invalid payload")
return [
{
"text": str(item.get("text", "")),
"start_ms": float(item.get("start_ms", 0.0)),
"end_ms": float(item.get("end_ms", 0.0)),
}
for item in items
if isinstance(item, dict) and str(item.get("text", "")).strip()
]
def create_stream(
self,
*,
chunk_size_sec: float = 1.2,
unfixed_chunk_num: int = 2,
rollback_tokens: int = 5,
max_new_tokens: int = 32,
language: Optional[str] = None,
) -> QwenASRRustStreamHandle:
handle = self._lib.qwen_asr_stream_new()
if not handle:
raise RuntimeError("QwenASR failed to create stream state")
self._configure_stream(
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=rollback_tokens,
max_new_tokens=max_new_tokens,
past_text=True,
)
self._set_language(language)
return QwenASRRustStreamHandle(self._lib, handle)
def push_stream(
self,
stream: QwenASRRustStreamHandle,
samples: np.ndarray,
*,
chunk_size_sec: float,
unfixed_chunk_num: int,
rollback_tokens: int,
max_new_tokens: int,
language: Optional[str],
finalize: bool = False,
) -> str:
self._configure_stream(
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=rollback_tokens,
max_new_tokens=max_new_tokens,
past_text=True,
)
self._set_language(language)
pcm = np.ascontiguousarray(samples, dtype=np.float32)
pointer = (
pcm.ctypes.data_as(ctypes.POINTER(ctypes.c_float))
if len(pcm) > 0
else None
)
delta_ptr = self._lib.qwen_asr_stream_push(
self._engine,
stream.handle,
pointer,
len(pcm),
1 if finalize else 0,
)
delta_text = _decode_and_free_string(self._lib, delta_ptr) or ""
stream.accumulated_text += delta_text
return stream.accumulated_text
class QwenASRRustRuntime:
"""Higher-level Rust runtime bundle for ASR + aligner + streaming."""
def __init__(
self,
model_path: str,
*,
forced_aligner_path: Optional[str] = None,
num_threads: int = 0,
verbosity: int = 0,
) -> None:
self._asr = QwenASRRustBackend(
model_path=model_path,
num_threads=num_threads,
verbosity=verbosity,
)
self._aligner: Optional[QwenASRRustBackend] = None
if forced_aligner_path:
self._aligner = QwenASRRustBackend(
model_path=forced_aligner_path,
num_threads=num_threads,
verbosity=verbosity,
)
def transcribe_file(self, audio_path: str, language: Optional[str] = None) -> str:
return self._asr.transcribe_file(audio_path=audio_path, language=language)
def align_transcript(
self,
audio_path: str,
text: str,
language: Optional[str] = None,
) -> list[dict[str, float | str]]:
transcript = text.strip()
if not transcript:
return []
if self._aligner is None:
raise RuntimeError("Forced alignment requires a configured aligner model")
return self._aligner.force_align_file(
audio_path=audio_path,
text=transcript,
language=guess_alignment_language(transcript, language),
)
def create_stream(
self,
*,
chunk_size_sec: float = 1.2,
unfixed_chunk_num: int = 2,
rollback_tokens: int = 5,
max_new_tokens: int = 32,
language: Optional[str] = None,
) -> QwenASRRustStreamHandle:
return self._asr.create_stream(
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=rollback_tokens,
max_new_tokens=max_new_tokens,
language=language,
)
def push_stream(
self,
stream: QwenASRRustStreamHandle,
samples: np.ndarray,
*,
chunk_size_sec: float,
unfixed_chunk_num: int,
rollback_tokens: int,
max_new_tokens: int,
language: Optional[str],
) -> str:
return self._asr.push_stream(
stream=stream,
samples=samples,
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=rollback_tokens,
max_new_tokens=max_new_tokens,
language=language,
finalize=False,
)
def finish_stream(
self,
stream: QwenASRRustStreamHandle,
*,
chunk_size_sec: float,
unfixed_chunk_num: int,
rollback_tokens: int,
max_new_tokens: int,
language: Optional[str],
) -> str:
return self._asr.push_stream(
stream=stream,
samples=np.array([], dtype=np.float32),
chunk_size_sec=chunk_size_sec,
unfixed_chunk_num=unfixed_chunk_num,
rollback_tokens=rollback_tokens,
max_new_tokens=max_new_tokens,
language=language,
finalize=True,
)

View File

@ -0,0 +1,16 @@
# -*- coding: utf-8 -*-
"""ASR runtime routing and pooling layer."""
from .router import (
OfflineASRRequest,
RuntimeEngineLease,
RuntimeRouter,
get_runtime_router,
)
__all__ = [
"OfflineASRRequest",
"RuntimeEngineLease",
"RuntimeRouter",
"get_runtime_router",
]

View File

@ -0,0 +1,48 @@
# -*- coding: utf-8 -*-
"""Small async pool for per-request ASR engines."""
from __future__ import annotations
import asyncio
import threading
from dataclasses import dataclass
from typing import Callable, Generic, Optional, TypeVar
T = TypeVar("T")
@dataclass
class _PoolState(Generic[T]):
queue: asyncio.Queue[T]
class LocalEnginePool(Generic[T]):
"""Fixed-size lazy engine pool backed by ``asyncio.Queue``."""
def __init__(self, size: int, factory: Callable[[], T]):
self._size = max(1, size)
self._factory = factory
self._state: Optional[_PoolState[T]] = None
self._init_lock = threading.Lock()
def _ensure_state(self) -> _PoolState[T]:
if self._state is not None:
return self._state
with self._init_lock:
if self._state is None:
self._state = _PoolState(queue=asyncio.Queue(maxsize=self._size))
for _ in range(self._size):
self._state.queue.put_nowait(self._factory())
return self._state
def warmup(self) -> None:
self._ensure_state()
async def acquire(self) -> T:
state = self._ensure_state()
return await state.queue.get()
async def release(self, engine: T) -> None:
state = self._ensure_state()
await state.queue.put(engine)

View File

@ -0,0 +1,241 @@
# -*- coding: utf-8 -*-
"""Runtime router for pooled ASR execution."""
from __future__ import annotations
import asyncio
import threading
from dataclasses import dataclass
from enum import Enum
from typing import Awaitable, Callable, Optional
import torch
from app.core.accelerator import get_accelerator_info
from app.core.config import settings
from app.core.device import detect_device
from app.core.executor import run_sync
from app.services.asr.engines import ASRFullResult, BaseASREngine
from app.services.asr.manager import get_model_manager
from app.services.asr.qwenasr_rust import is_qwenasr_rust_available
from .local_pool import LocalEnginePool
class RuntimeFamily(str, Enum):
QWEN_VLLM = "qwen_vllm"
QWEN_RUST_CPU = "qwen_rust_cpu"
@dataclass
class OfflineASRRequest:
model_id: str
audio_path: str
hotwords: str = ""
enable_punctuation: bool = True
enable_itn: bool = True
sample_rate: int = 16000
enable_speaker_diarization: bool = True
enable_speaker_identification: bool = True
enable_text_cleanup: bool = True
word_timestamps: bool = False
timestamp_scale: float = 1.0
task_id: Optional[str] = None
progress_callback: Optional[Callable[[str, str, int, Optional[dict[str, object]]], None]] = None
class RuntimeEngineLease:
"""Lifecycle wrapper around a pooled engine instance."""
def __init__(self, engine: BaseASREngine, release_callback: Callable[[], None | Awaitable[None]]):
self.engine = engine
self._release_callback = release_callback
self._closed = False
async def close(self) -> None:
if self._closed:
return
self._closed = True
result = self._release_callback()
if asyncio.iscoroutine(result):
await result
async def __aenter__(self) -> BaseASREngine:
return self.engine
async def __aexit__(self, exc_type, exc, tb) -> None:
await self.close()
class RuntimeRouter:
"""Central backend router for all ASR entrypoints."""
def __init__(self):
self._manager = get_model_manager()
self._pools: dict[tuple[RuntimeFamily, str], LocalEnginePool[BaseASREngine]] = {}
self._shared_engines: dict[tuple[RuntimeFamily, str], BaseASREngine] = {}
self._shared_limits: dict[tuple[RuntimeFamily, str], asyncio.Semaphore] = {}
self._pool_lock = threading.Lock()
self._loaded_model_ids: set[str] = set()
def resolve_model_id(self, model_id: Optional[str]) -> str:
if model_id:
return model_id
config = self._manager.get_declared_entry_config()
return config.model_id
def _resolve_family(self, model_id: str) -> RuntimeFamily:
device = detect_device(settings.DEVICE)
accelerator = get_accelerator_info()
if model_id.startswith("qwen3-asr-"):
if accelerator.is_gpu and device.startswith("cuda"):
return RuntimeFamily.QWEN_VLLM
if device == "cpu" and is_qwenasr_rust_available():
return RuntimeFamily.QWEN_RUST_CPU
raise RuntimeError(
"Qwen3-ASR is not available on "
f"accelerator='{accelerator.vendor}' device='{device}'"
)
raise RuntimeError(f"Unsupported runtime model: {model_id}")
def _pool_size_for_family(self, family: RuntimeFamily) -> int:
if family == RuntimeFamily.QWEN_VLLM:
return 1
return settings.QWEN_RUST_CPU_WORKERS
def _create_pool(self, family: RuntimeFamily, model_id: str) -> LocalEnginePool[BaseASREngine]:
pool_key = (family, model_id)
existing = self._pools.get(pool_key)
if existing is not None:
return existing
with self._pool_lock:
existing = self._pools.get(pool_key)
if existing is not None:
return existing
pool = LocalEnginePool(
size=self._pool_size_for_family(family),
factory=lambda: self._manager.create_engine(model_id),
)
self._pools[pool_key] = pool
self._loaded_model_ids.add(model_id)
return pool
def _get_shared_engine(self, family: RuntimeFamily, model_id: str) -> tuple[BaseASREngine, asyncio.Semaphore]:
runtime_key = (family, model_id)
engine = self._shared_engines.get(runtime_key)
semaphore = self._shared_limits.get(runtime_key)
if engine is not None and semaphore is not None:
return engine, semaphore
with self._pool_lock:
engine = self._shared_engines.get(runtime_key)
semaphore = self._shared_limits.get(runtime_key)
if engine is None:
engine = self._manager.create_engine(model_id)
self._shared_engines[runtime_key] = engine
self._loaded_model_ids.add(model_id)
if semaphore is None:
shared_concurrency = max(1, int(settings.QWEN_VLLM_SHARED_CONCURRENCY))
semaphore = asyncio.Semaphore(shared_concurrency)
self._shared_limits[runtime_key] = semaphore
return engine, semaphore
def warmup_model(self, model_id: Optional[str] = None) -> None:
resolved_model_id = self.resolve_model_id(model_id)
family = self._resolve_family(resolved_model_id)
if family == RuntimeFamily.QWEN_VLLM:
self._get_shared_engine(family, resolved_model_id)
return
pool = self._create_pool(family, resolved_model_id)
pool.warmup()
def get_loaded_model_ids(self) -> list[str]:
return sorted(self._loaded_model_ids)
def get_memory_usage(self) -> dict[str, object]:
memory_info: dict[str, object] = {
"model_list": self.get_loaded_model_ids(),
"loaded_count": len(self._loaded_model_ids),
}
accelerator = get_accelerator_info()
memory_info["accelerator"] = accelerator.as_dict()
if accelerator.is_gpu and torch.cuda.is_available():
memory_info["gpu_memory"] = {
"allocated": f"{torch.cuda.memory_allocated() / 1024**3:.2f}GB",
"cached": f"{torch.cuda.memory_reserved() / 1024**3:.2f}GB",
"max_allocated": f"{torch.cuda.max_memory_allocated() / 1024**3:.2f}GB",
}
return memory_info
async def acquire_engine(self, model_id: Optional[str] = None) -> RuntimeEngineLease:
resolved_model_id = self.resolve_model_id(model_id)
family = self._resolve_family(resolved_model_id)
if family == RuntimeFamily.QWEN_VLLM:
engine, semaphore = self._get_shared_engine(family, resolved_model_id)
await semaphore.acquire()
return RuntimeEngineLease(
engine=engine,
release_callback=semaphore.release,
)
pool = self._create_pool(family, resolved_model_id)
engine = await pool.acquire()
return RuntimeEngineLease(
engine=engine,
release_callback=lambda: pool.release(engine),
)
async def run_offline(self, request: OfflineASRRequest) -> ASRFullResult:
async with await self.acquire_engine(request.model_id) as engine:
result = await run_sync(
engine.transcribe_long_audio,
audio_path=request.audio_path,
hotwords=request.hotwords,
enable_punctuation=request.enable_punctuation,
enable_itn=request.enable_itn,
sample_rate=request.sample_rate,
enable_speaker_diarization=request.enable_speaker_diarization,
enable_speaker_identification=request.enable_speaker_identification,
enable_text_cleanup=request.enable_text_cleanup,
word_timestamps=request.word_timestamps,
timestamp_scale=request.timestamp_scale,
task_id=request.task_id,
progress_callback=request.progress_callback,
)
if (
request.enable_speaker_diarization
and request.enable_speaker_identification
and settings.SPEAKER_DB_ENABLED
):
try:
from app.services.speaker_registry import get_speaker_registry_service
if request.progress_callback is not None:
request.progress_callback(
"speaker_matching",
"正在匹配已注册声纹库。",
94,
None,
)
result = await get_speaker_registry_service().apply_registered_speakers(
result,
threshold=settings.SV_THRESHOLD,
)
except Exception:
# 数据库/声纹匹配失败不影响原有 ASR 输出,仍保留 CAM++ 的局部说话人编号。
pass
return result
_runtime_router: Optional[RuntimeRouter] = None
_runtime_router_lock = threading.Lock()
def get_runtime_router() -> RuntimeRouter:
global _runtime_router
if _runtime_router is None:
with _runtime_router_lock:
if _runtime_router is None:
_runtime_router = RuntimeRouter()
return _runtime_router

View File

@ -0,0 +1,10 @@
# -*- coding: utf-8 -*-
"""
音频处理服务模块
提供统一的音频处理服务层,封装音频下载、格式转换、归一化等功能。
"""
from .audio_service import AudioProcessingService, get_audio_service
__all__ = ["AudioProcessingService", "get_audio_service"]

View File

@ -0,0 +1,234 @@
# -*- coding: utf-8 -*-
"""
音频处理服务
封装音频处理逻辑,提供统一的音频下载、格式转换、归一化等服务。
API层应该通过此服务层处理音频,而不是直接调用 utils/audio.py 中的函数。
"""
import logging
import threading
from dataclasses import dataclass
from typing import Optional
from fastapi import Request
from ...core.config import settings
from ...core.exceptions import InvalidMessageException
from ...utils.audio import (
download_audio_from_url,
save_audio_to_temp_file,
normalize_audio_for_asr,
get_audio_duration,
cleanup_temp_file,
get_audio_file_suffix,
)
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class AudioProcessingResult:
normalized_path: str
duration: float
original_path: str
timestamp_scale: float = 1.0
class AudioProcessingService:
"""音频处理服务
提供统一的音频处理接口,包括:
1. 从URL下载音频
2. 处理上传的音频文件
3. 音频格式转换和归一化
4. 临时文件管理
"""
async def process_from_request(
self,
request: Request,
audio_address: Optional[str] = None,
task_id: Optional[str] = None,
sample_rate: Optional[int] = None,
) -> AudioProcessingResult:
"""从请求中处理音频
支持两种方式:
1. 请求体上传:从请求体读取二进制音频/视频数据
2. URL下载:通过 audio_address 参数指定音频/视频 URL
当请求体和 audio_address 同时存在时,优先使用请求体,
并忽略 audio_address。
Args:
request: FastAPI请求对象
audio_address: 音频文件URL(可选)
task_id: 任务ID,用于日志记录(可选)
sample_rate: 目标采样率(可选,默认16000)
Returns:
Processed audio path, duration, original path, and timestamp metadata.
Raises:
InvalidMessageException: 音频数据为空或文件太大
InvalidParameterException: URL无效或下载失败
"""
task_id = task_id or "unknown"
target_sr = sample_rate or 16000
# 优先读取请求体;若请求体为空,再回退到 audio_address。
# 注意:对于 FastAPI 已经解析过 form/multipart 的请求,
# 再次读取 body 可能抛出 "Stream consumed"。
try:
uploaded_data = await request.body()
except RuntimeError as exc:
if "Stream consumed" in str(exc):
logger.info(f"[{task_id}] 请求体已被上游读取,跳过 request.body() 回退逻辑")
uploaded_data = b""
else:
raise
if uploaded_data:
if audio_address:
logger.info(f"[{task_id}] 检测到同时提供上传内容和 audio_address,已忽略 audio_address")
return self._process_audio_bytes(
audio_data=uploaded_data,
filename=None,
task_id=task_id,
target_sr=target_sr,
)
if audio_address:
logger.info(f"[{task_id}] 开始从URL下载音频: {audio_address}")
audio_data = download_audio_from_url(audio_address)
logger.info(
f"[{task_id}] 音频下载完成,大小: {len(audio_data) / 1024 / 1024:.2f}MB"
)
return self._process_audio_bytes(
audio_data=audio_data,
filename=audio_address,
task_id=task_id,
target_sr=target_sr,
)
raise InvalidMessageException("音频数据为空", task_id)
async def process_upload_file(
self,
audio_data: bytes,
filename: Optional[str] = None,
task_id: Optional[str] = None,
sample_rate: Optional[int] = None,
) -> AudioProcessingResult:
"""处理上传的音频文件
Args:
audio_data: 音频二进制数据
filename: 原始文件名(用于检测格式,可选)
task_id: 任务ID,用于日志记录(可选)
sample_rate: 目标采样率(可选,默认16000)
Returns:
Processed audio path, duration, original path, and timestamp metadata.
Raises:
InvalidMessageException: 音频数据为空或文件太大
"""
task_id = task_id or "unknown"
target_sr = sample_rate or 16000
return self._process_audio_bytes(
audio_data=audio_data,
filename=filename,
task_id=task_id,
target_sr=target_sr,
)
def _process_audio_bytes(
self,
*,
audio_data: bytes,
filename: Optional[str],
task_id: str,
target_sr: int,
) -> AudioProcessingResult:
"""Persist, normalize, and measure audio bytes."""
audio_path = None
normalized_audio_path = None
try:
if not audio_data:
raise InvalidMessageException("音频数据为空", task_id)
file_size = len(audio_data)
logger.info(f"[{task_id}] 音频文件大小: {file_size / 1024 / 1024:.2f}MB")
# 检查文件大小
if file_size > settings.MAX_AUDIO_SIZE:
max_mb = settings.MAX_AUDIO_SIZE // 1024 // 1024
raise InvalidMessageException(
f"音频文件太大,最大支持{max_mb}MB", task_id
)
file_suffix = get_audio_file_suffix(
audio_address=filename,
audio_data=audio_data,
)
logger.info(f"[{task_id}] 识别文件格式: {file_suffix}")
audio_path = save_audio_to_temp_file(audio_data, file_suffix)
logger.info(f"[{task_id}] 临时文件: {audio_path}")
logger.info(f"[{task_id}] 开始音频格式转换...")
normalized_audio = normalize_audio_for_asr(audio_path, target_sr)
normalized_audio_path = normalized_audio.path
logger.info(f"[{task_id}] 音频格式转换完成: {normalized_audio_path}")
decoded_duration = get_audio_duration(normalized_audio_path)
audio_duration = decoded_duration * normalized_audio.timestamp_scale
logger.info(f"[{task_id}] 音频时长: {audio_duration:.1f}s")
return AudioProcessingResult(
normalized_path=normalized_audio_path,
duration=audio_duration,
original_path=audio_path,
timestamp_scale=normalized_audio.timestamp_scale,
)
except Exception:
if audio_path:
cleanup_temp_file(audio_path)
if normalized_audio_path and normalized_audio_path != audio_path:
cleanup_temp_file(normalized_audio_path)
raise
def cleanup(
self, audio_path: Optional[str], normalized_path: Optional[str] = None
) -> None:
"""清理临时文件
Args:
audio_path: 原始音频文件路径
normalized_path: 归一化后的音频文件路径(可选)
"""
if audio_path:
cleanup_temp_file(audio_path)
if normalized_path and normalized_path != audio_path:
cleanup_temp_file(normalized_path)
# 全局服务实例(单例模式)
_audio_service: Optional[AudioProcessingService] = None
_audio_service_lock = threading.Lock()
def get_audio_service() -> AudioProcessingService:
"""获取音频处理服务实例(线程安全的单例)
Returns:
AudioProcessingService: 音频处理服务实例
"""
global _audio_service
if _audio_service is None:
with _audio_service_lock:
if _audio_service is None:
_audio_service = AudioProcessingService()
return _audio_service

File diff suppressed because it is too large Load Diff

View File

@ -0,0 +1,640 @@
# -*- coding: utf-8 -*-
"""FunASR-style realtime speaker chunking and clustering."""
from __future__ import annotations
import logging
import re
from collections import defaultdict
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import scipy.linalg
import sklearn.metrics.pairwise
from sklearn.cluster._kmeans import k_means
from app.core.config import settings
from app.core.executor import run_sync
from app.services.speaker_registry import get_speaker_registry_service
logger = logging.getLogger(__name__)
class RealtimeSpeakerClusterer:
"""Sentence-level speaker attribution using chunk embeddings and session clustering."""
FUNASR_SMOOTH_MINDUR_SEC = 0.7
_FAST_ATTACH_MAX_SEC = 1.6
_REGISTRY_MATCH_MIN_SEC = 2.4
def _is_generic_speaker_name(self, value: Any) -> bool:
text = str(value or "").strip()
return bool(re.fullmatch(r"Speaker\d+", text))
def _has_named_identity(self, record: Optional[Dict[str, Any]]) -> bool:
if not record:
return False
return bool(
record.get("registry_speaker_id")
or record.get("user_id")
or (
record.get("speaker_name")
and not self._is_generic_speaker_name(record.get("speaker_name"))
)
)
def _normalize_embedding(self, embedding: np.ndarray) -> np.ndarray:
emb = np.asarray(embedding, dtype=np.float32).reshape(-1)
norm = max(float(np.linalg.norm(emb)), 1e-12)
return (emb / norm).astype(np.float32)
@staticmethod
def _coerce_cluster_index(value: Any, fallback: int) -> int:
try:
if value is None:
return int(fallback)
return int(value)
except (TypeError, ValueError):
return int(fallback)
def _match_existing_speaker(
self,
speaker_records: List[Dict[str, Any]],
mean_embedding: np.ndarray,
) -> Optional[Dict[str, Any]]:
if not speaker_records:
return None
best_record: Optional[Dict[str, Any]] = None
best_score = -1.0
for record in speaker_records:
record_embedding = record.get("embedding")
if record_embedding is None:
continue
score = float(np.dot(self._normalize_embedding(record_embedding), mean_embedding))
if score > best_score:
best_score = score
best_record = record
if best_record is None:
return None
base_threshold = max(
float(getattr(settings, "REALTIME_UNKNOWN_SPK_CLUSTER_THRESHOLD", 0.58) or 0.58),
0.3,
)
has_named_identity = bool(
best_record.get("registry_speaker_id")
or best_record.get("user_id")
or (
best_record.get("speaker_name")
and not self._is_generic_speaker_name(best_record.get("speaker_name"))
)
)
threshold = max(base_threshold, 0.62 if has_named_identity else 0.72)
if best_score < threshold:
return None
matched = dict(best_record)
matched["_match_score"] = best_score
return matched
def _next_generic_speaker_id(
self,
speaker_records: List[Dict[str, Any]],
) -> str:
seen: set[int] = set()
for record in speaker_records:
for value in (
record.get("speaker_id"),
record.get("speaker_name"),
):
text = str(value or "").strip()
match = re.fullmatch(r"Speaker(\d+)", text)
if match:
seen.add(int(match.group(1)))
next_index = max(seen, default=0) + 1
return f"Speaker{next_index:02d}"
def _is_mixed_speaker_segment(
self,
speaker_records: List[Dict[str, Any]],
current_chunks: List[Dict[str, Any]],
) -> bool:
if len(current_chunks) < 2 or len(speaker_records) < 2:
return False
assigned_labels: List[str] = []
for chunk in current_chunks:
embedding = chunk.get("embedding")
if embedding is None:
continue
matched = self._match_existing_speaker(
speaker_records,
self._normalize_embedding(np.asarray(embedding, dtype=np.float32)),
)
if not matched:
continue
label = str(
matched.get("registry_speaker_id")
or matched.get("speaker_name")
or matched.get("speaker_id")
or ""
).strip()
if label:
assigned_labels.append(label)
if len(assigned_labels) < 2:
return False
return len(set(assigned_labels)) >= 2
def _correct_labels(self, labels: np.ndarray) -> np.ndarray:
labels_id = 0
id2id: dict[int, int] = {}
new_labels: list[int] = []
for label in labels.tolist():
label = int(label)
if label not in id2id:
id2id[label] = labels_id
labels_id += 1
new_labels.append(id2id[label])
return np.asarray(new_labels, dtype=np.int32)
def _spectral_cluster(self, X: np.ndarray, oracle_num: Optional[int] = None) -> np.ndarray:
sim_mat = sklearn.metrics.pairwise.cosine_similarity(X, X)
A = sim_mat.copy()
pval = 0.022
if A.shape[0] * pval < 6:
pval = 6.0 / A.shape[0]
n_elems = int((1 - pval) * A.shape[0])
for i in range(A.shape[0]):
low_indexes = np.argsort(A[i, :])[:n_elems]
A[i, low_indexes] = 0
A = 0.5 * (A + A.T)
A[np.diag_indices(A.shape[0])] = 0
D = np.diag(np.sum(np.abs(A), axis=1))
L = D - A
lambdas, eig_vecs = scipy.linalg.eigh(L)
if oracle_num is not None:
num_spk = max(int(oracle_num), 1)
else:
min_num_spks = 1
max_num_spks = min(15, max(1, X.shape[0] - 1))
eig_slice = lambdas[min_num_spks - 1 : max_num_spks + 1]
gap_list = [float(eig_slice[i + 1]) - float(eig_slice[i]) for i in range(len(eig_slice) - 1)]
num_spk = int(np.argmax(gap_list)) + min_num_spks if gap_list else 1
emb = eig_vecs[:, : max(num_spk, 1)]
_, labels, _ = k_means(emb, max(num_spk, 1))
return np.asarray(labels, dtype=np.int32)
def _merge_by_cos(self, labels: np.ndarray, embs: np.ndarray, cos_thr: float) -> np.ndarray:
labels = np.asarray(labels, dtype=np.int32).copy()
while True:
spk_num = int(labels.max()) + 1
if spk_num <= 1:
break
centers = []
for i in range(spk_num):
spk_emb = embs[labels == i].mean(0)
centers.append(spk_emb)
centers = np.stack(centers, axis=0)
norm_centers = centers / np.linalg.norm(centers, axis=1, keepdims=True)
affinity = np.matmul(norm_centers, norm_centers.T)
affinity = np.triu(affinity, 1)
spks = np.unravel_index(np.argmax(affinity), affinity.shape)
if float(affinity[spks]) < cos_thr:
break
for i in range(len(labels)):
if labels[i] == spks[1]:
labels[i] = spks[0]
elif labels[i] > spks[1]:
labels[i] -= 1
return self._correct_labels(labels)
def _cluster_embeddings(self, X: np.ndarray, oracle_num: Optional[int] = None) -> np.ndarray:
if X.shape[0] < 20:
labels = np.zeros(X.shape[0], dtype=np.int32)
else:
labels = self._spectral_cluster(X, oracle_num=oracle_num)
return self._merge_by_cos(labels, X, cos_thr=0.78 if oracle_num is None else 1.0)
def build_sv_chunks(
self,
audio: np.ndarray,
*,
segment_start_ms: int = 0,
) -> List[Dict[str, Any]]:
audio = np.asarray(audio, dtype=np.float32).flatten()
if audio.size == 0:
return []
sample_rate = 16000
chunk_len = int(1.5 * sample_rate)
chunk_shift = int(0.75 * sample_rate)
chunks: List[Dict[str, Any]] = []
last_chunk_end = 0
for chunk_start in range(0, audio.shape[0], chunk_shift):
chunk_end = min(chunk_start + chunk_len, audio.shape[0])
if chunk_end <= last_chunk_end:
break
actual_start = max(0, chunk_end - chunk_len)
actual_end = chunk_end
last_chunk_end = actual_end
chunk_audio = np.asarray(audio[actual_start:actual_end], dtype=np.float32)
if chunk_audio.shape[0] < chunk_len:
chunk_audio = np.pad(chunk_audio, (0, chunk_len - chunk_audio.shape[0]), "constant")
chunks.append(
{
"start_ms": int(segment_start_ms + actual_start / 16.0),
"end_ms": int(segment_start_ms + actual_end / 16.0),
"audio": chunk_audio,
}
)
return chunks
async def extract_chunk_embeddings(
self,
audio: np.ndarray,
*,
segment_start_ms: int = 0,
) -> List[Dict[str, Any]]:
chunk_items = self.build_sv_chunks(audio, segment_start_ms=segment_start_ms)
if not chunk_items:
chunk_items = [
{
"start_ms": int(segment_start_ms),
"end_ms": int(segment_start_ms + np.asarray(audio, dtype=np.float32).size / 16.0),
"audio": np.asarray(audio, dtype=np.float32),
}
]
registry = get_speaker_registry_service()
results: List[Dict[str, Any]] = []
for item in chunk_items:
embedding = await run_sync(
registry.extract_embedding_from_audio,
item["audio"],
16000,
model_id=settings.REALTIME_SV_MODEL,
model_revision=settings.REALTIME_SV_MODEL_REVISION or None,
)
results.append(
{
"start_ms": int(item["start_ms"]),
"end_ms": int(item["end_ms"]),
"embedding": self._normalize_embedding(np.asarray(embedding, dtype=np.float32)),
}
)
return results
async def extract_registry_embedding(self, audio: np.ndarray) -> Optional[np.ndarray]:
audio_array = np.asarray(audio, dtype=np.float32).flatten()
if audio_array.size == 0:
return None
try:
registry = get_speaker_registry_service()
embedding = await run_sync(
registry.extract_embedding_from_audio,
audio_array,
16000,
)
return self._normalize_embedding(np.asarray(embedding, dtype=np.float32))
except Exception as exc:
logger.warning("Realtime registry speaker embedding failed: %s", exc)
return None
def _flatten_chunk_records(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], np.ndarray]:
flat_chunks: List[Dict[str, Any]] = []
flat_embeddings: List[np.ndarray] = []
for record_idx, record in enumerate(records):
chunk_entries = record.get("chunks") or []
if not chunk_entries:
embeddings = record.get("embeddings") or []
if not embeddings and record.get("embedding") is not None:
embeddings = [record["embedding"]]
start_ms = int(record.get("start_ms", 0))
end_ms = int(record.get("end_ms", start_ms))
for embedding in embeddings:
chunk_entries.append(
{
"start_ms": start_ms,
"end_ms": end_ms,
"embedding": embedding,
}
)
for chunk in chunk_entries:
emb = self._normalize_embedding(np.asarray(chunk["embedding"], dtype=np.float32))
flat_chunks.append(
{
"record_index": record_idx,
"start_ms": int(chunk.get("start_ms", record.get("start_ms", 0))),
"end_ms": int(chunk.get("end_ms", record.get("end_ms", 0))),
}
)
flat_embeddings.append(emb)
if not flat_embeddings:
return [], np.zeros((0, 0), dtype=np.float32)
return flat_chunks, np.stack(flat_embeddings, axis=0).astype(np.float32)
def _merge_seque(self, distribute_res: List[List[float]]) -> List[List[float]]:
if not distribute_res:
return []
res = [distribute_res[0][:]]
for item in distribute_res[1:]:
if item[2] != res[-1][2] or item[0] > res[-1][1]:
res.append(item[:])
else:
res[-1][1] = item[1]
return res
def _smooth_timeline(
self,
res: List[List[float]],
mindur: float = FUNASR_SMOOTH_MINDUR_SEC,
) -> List[List[float]]:
if len(res) < 2:
return res
for item in res:
item[0] = round(float(item[0]), 2)
item[1] = round(float(item[1]), 2)
for idx in range(len(res)):
if res[idx][1] - res[idx][0] < mindur:
if idx == 0:
res[idx][2] = res[idx + 1][2]
elif idx == len(res) - 1:
res[idx][2] = res[idx - 1][2]
elif res[idx][0] - res[idx - 1][1] <= res[idx + 1][0] - res[idx][1]:
res[idx][2] = res[idx - 1][2]
else:
res[idx][2] = res[idx + 1][2]
return self._merge_seque(res)
def _postprocess_timeline(
self,
flat_chunks: List[Dict[str, Any]],
labels: np.ndarray,
embeddings: np.ndarray,
) -> List[Dict[str, Any]]:
assert len(flat_chunks) == len(labels)
labels = self._correct_labels(labels)
distribute_res = [
[
float(chunk["start_ms"]) / 1000.0,
float(chunk["end_ms"]) / 1000.0,
int(labels[idx]),
]
for idx, chunk in enumerate(flat_chunks)
]
distribute_res = self._merge_seque(distribute_res)
def is_overlapped(t1: float, t2: float) -> bool:
return t1 > t2 + 1e-4
for idx in range(1, len(distribute_res)):
if is_overlapped(distribute_res[idx - 1][1], distribute_res[idx][0]):
pivot = (distribute_res[idx][0] + distribute_res[idx - 1][1]) / 2.0
distribute_res[idx][0] = pivot
distribute_res[idx - 1][1] = pivot
distribute_res = self._smooth_timeline(distribute_res)
return [
{
"start_ms": int(round(item[0] * 1000.0)),
"end_ms": int(round(item[1] * 1000.0)),
"cluster_index": int(item[2]),
}
for item in distribute_res
]
def _pick_segment_cluster(
self,
cluster_ranges: List[Dict[str, Any]],
*,
segment_start_ms: int,
segment_end_ms: int,
) -> int:
overlaps: Dict[int, int] = defaultdict(int)
for item in cluster_ranges:
overlap = min(segment_end_ms, int(item["end_ms"])) - max(segment_start_ms, int(item["start_ms"]))
if overlap > 0:
overlaps[int(item["cluster_index"])] += int(overlap)
if overlaps:
return max(overlaps.items(), key=lambda kv: (kv[1], -kv[0]))[0]
centers = [
item for item in cluster_ranges
if int(item["start_ms"]) <= segment_end_ms and int(item["end_ms"]) >= segment_start_ms
]
if centers:
return int(centers[0]["cluster_index"])
return 0
def cluster_records_with_ranges(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], List[int], List[Dict[str, Any]]]:
flat_chunks, X = self._flatten_chunk_records(records)
if X.size == 0:
return [], [], []
labels = self._cluster_embeddings(X, oracle_num=None)
cluster_ranges = self._postprocess_timeline(flat_chunks, labels, X)
cluster_ids = sorted({int(item["cluster_index"]) for item in cluster_ranges})
if not cluster_ids:
cluster_ids = sorted({int(label) for label in labels.tolist()})
clusters: List[Dict[str, Any]] = []
for cluster_idx in cluster_ids:
member_mask = labels == int(cluster_idx)
member_embeddings = X[member_mask]
centroid = self._normalize_embedding(member_embeddings.mean(0))
record_refs = []
seen_record_indices = set()
for emb_idx, chunk in enumerate(flat_chunks):
record_idx = int(chunk["record_index"])
if not member_mask[emb_idx] or record_idx in seen_record_indices:
continue
seen_record_indices.add(record_idx)
record_refs.append(records[record_idx])
clusters.append(
{
"centroid": centroid,
"count": int(member_mask.sum()),
"record_refs": record_refs,
}
)
record_to_labels: List[List[int]] = [[] for _ in records]
for emb_idx, chunk in enumerate(flat_chunks):
record_to_labels[int(chunk["record_index"])].append(int(labels[emb_idx]))
record_assignments: List[int] = []
for record, local_assignments in zip(records, record_to_labels):
if local_assignments:
cluster_idx = self._pick_segment_cluster(
cluster_ranges,
segment_start_ms=int(record.get("start_ms", 0)),
segment_end_ms=int(record.get("end_ms", record.get("start_ms", 0))),
)
record_assignments.append(cluster_idx)
else:
record_assignments.append(-1)
return clusters, record_assignments, cluster_ranges
def cluster_records(
self,
records: List[Dict[str, Any]],
) -> Tuple[List[Dict[str, Any]], List[int]]:
clusters, record_assignments, _ = self.cluster_records_with_ranges(records)
return clusters, record_assignments
async def resolve_segment_speaker(
self,
speaker_records: List[Dict[str, Any]],
audio: np.ndarray,
*,
segment_start_ms: int,
segment_end_ms: int,
enable_registry_match: bool,
speaker_threshold: Optional[float],
) -> Optional[Dict[str, Any]]:
duration_sec = float(len(audio)) / 16000.0
last_record = speaker_records[-1] if speaker_records else None
if duration_sec < self._FAST_ATTACH_MAX_SEC and speaker_records:
if self._has_named_identity(last_record):
speaker_id = last_record.get("registry_speaker_id") or last_record.get("speaker_id") or "Speaker01"
speaker_name = last_record.get("speaker_name") or speaker_id
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": last_record.get("user_id"),
"registry_speaker_id": last_record.get("registry_speaker_id"),
"speaker_confidence": 0.0,
"speaker_strategy": "short_attach",
"_cluster_index": last_record.get("cluster_index"),
"_embedding": last_record.get("embedding"),
"_chunk_embeddings": last_record.get("embeddings") or [],
"_chunks": last_record.get("chunks") or [],
}
current_chunks = await self.extract_chunk_embeddings(
np.asarray(audio, dtype=np.float32),
segment_start_ms=segment_start_ms,
)
if not current_chunks:
if speaker_records:
if self._has_named_identity(last_record):
speaker_id = last_record.get("registry_speaker_id") or last_record.get("speaker_id") or "Speaker01"
speaker_name = last_record.get("speaker_name") or speaker_id
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": last_record.get("user_id"),
"registry_speaker_id": last_record.get("registry_speaker_id"),
"speaker_confidence": 0.0,
"speaker_strategy": "embedding_attach",
"_cluster_index": last_record.get("cluster_index"),
"_embedding": last_record.get("embedding"),
"_chunk_embeddings": last_record.get("embeddings") or [],
"_chunks": last_record.get("chunks") or [],
}
return None
mean_embedding = self._normalize_embedding(
np.mean(np.stack([chunk["embedding"] for chunk in current_chunks], axis=0), axis=0)
)
if duration_sec >= 4.0 and self._is_mixed_speaker_segment(speaker_records, current_chunks):
return {
"speaker_id": -1,
"speaker_name": "",
"user_id": None,
"registry_speaker_id": None,
"speaker_confidence": 0.0,
"speaker_strategy": "mixed_segment",
"_cluster_index": None,
"_embedding": mean_embedding,
"_chunk_embeddings": [np.asarray(chunk["embedding"], dtype=np.float32) for chunk in current_chunks],
"_chunks": current_chunks,
}
matched_record = self._match_existing_speaker(speaker_records, mean_embedding)
speaker_id = self._next_generic_speaker_id(speaker_records)
speaker_name = speaker_id
user_id = None
registry_speaker_id = None
strategy = "new_speaker"
cluster_index = max(
[
int(record.get("cluster_index", -1))
for record in speaker_records
if record.get("cluster_index") is not None
] or [-1]
) + 1
confidence = 0.0
if matched_record is not None:
speaker_id = (
matched_record.get("registry_speaker_id")
or matched_record.get("speaker_id")
or speaker_id
)
speaker_name = matched_record.get("speaker_name") or speaker_id
user_id = matched_record.get("user_id")
registry_speaker_id = matched_record.get("registry_speaker_id")
cluster_index = self._coerce_cluster_index(
matched_record.get("cluster_index"),
cluster_index,
)
confidence = float(matched_record.get("_match_score", 0.0))
strategy = "embedding_match"
if (
not registry_speaker_id
and enable_registry_match
and duration_sec >= self._REGISTRY_MATCH_MIN_SEC
and (
matched_record is None
or confidence >= max(float(getattr(settings, "REALTIME_SPEAKER_CONFIRM_THRESHOLD", 0.62) or 0.62), 0.7)
)
):
# Realtime clustering can use a dedicated fast model, while the
# registry stores embeddings from the registration model. Match the
# registry in its own embedding space instead of comparing vectors
# extracted by a different model.
registry_embedding = await self.extract_registry_embedding(audio)
if registry_embedding is not None:
matched = await get_speaker_registry_service().identify_embedding(
registry_embedding,
threshold=speaker_threshold,
)
if matched.get("name"):
speaker_id = matched.get("speaker_id") or speaker_id
speaker_name = matched.get("name") or speaker_name
user_id = matched.get("user_id")
registry_speaker_id = matched.get("speaker_id")
strategy = "registry_match"
return {
"speaker_id": speaker_id,
"speaker_name": speaker_name,
"user_id": user_id,
"registry_speaker_id": registry_speaker_id,
"speaker_confidence": round(confidence, 4),
"speaker_strategy": strategy,
"_cluster_index": self._coerce_cluster_index(cluster_index, 0),
"_embedding": mean_embedding,
"_chunk_embeddings": [np.asarray(chunk["embedding"], dtype=np.float32) for chunk in current_chunks],
"_chunks": current_chunks,
"_matched_existing": matched_record is not None,
}
_realtime_speaker_clusterer: Optional[RealtimeSpeakerClusterer] = None
def get_realtime_speaker_clusterer() -> RealtimeSpeakerClusterer:
global _realtime_speaker_clusterer
if _realtime_speaker_clusterer is None:
_realtime_speaker_clusterer = RealtimeSpeakerClusterer()
return _realtime_speaker_clusterer

View File

@ -0,0 +1,250 @@
# -*- coding: utf-8 -*-
"""Speaker embedding extraction, registration, and database identification."""
from __future__ import annotations
import asyncio
import logging
import os
import tempfile
import threading
from typing import Any, Optional
import librosa
import numpy as np
import soundfile as sf
import torch
from app.core.config import settings
from app.core.database import pg_speaker_db
logger = logging.getLogger(__name__)
def _normalize_embedding(embedding: np.ndarray) -> np.ndarray:
array = np.asarray(embedding, dtype=np.float32).flatten()
norm = float(np.linalg.norm(array))
if norm < 1e-12:
return array
return array / norm
class SpeakerRegistryService:
"""Owns speaker embedding models and pgvector lookups."""
def __init__(self) -> None:
self._pipelines: dict[tuple[str, str], Any] = {}
self._lock = threading.Lock()
self._inference_lock = threading.BoundedSemaphore(1)
def _device(self) -> str:
from app.core.device import detect_device
return detect_device(settings.DEVICE)
def _get_pipeline(
self,
model_id: Optional[str] = None,
model_revision: Optional[str] = None,
) -> Any:
effective_model_id = model_id or settings.SV_MODEL
effective_revision = model_revision if model_revision is not None else settings.SV_MODEL_REVISION
cache_key = (effective_model_id, effective_revision or "")
cached = self._pipelines.get(cache_key)
if cached is not None:
return cached
with self._lock:
cached = self._pipelines.get(cache_key)
if cached is not None:
return cached
from modelscope.pipelines import pipeline
from modelscope.utils.constant import Tasks
from app.infrastructure.model_utils import resolve_model_path
model_path = resolve_model_path(effective_model_id)
device = self._device()
kwargs: dict[str, Any] = {
"task": Tasks.speaker_verification,
"model": model_path,
"device": device,
}
if effective_revision:
kwargs["model_revision"] = effective_revision
logger.info("正在加载声纹识别模型: %s, device=%s", model_path, device)
pipeline_instance = pipeline(**kwargs)
if hasattr(pipeline_instance, "device_name"):
pipeline_instance.device_name = device
model = getattr(pipeline_instance, "model", None)
if model is not None and hasattr(model, "to"):
pipeline_instance.model = model.to(device)
logger.info("声纹识别模型加载成功")
self._pipelines[cache_key] = pipeline_instance
return pipeline_instance
def ensure_loaded(self) -> None:
"""Eagerly initialize the speaker verification pipeline at startup."""
self._get_pipeline()
def extract_embedding_from_audio(
self,
audio_data: np.ndarray,
sample_rate: int = 16000,
*,
model_id: Optional[str] = None,
model_revision: Optional[str] = None,
) -> np.ndarray:
audio = np.asarray(audio_data, dtype=np.float32).flatten()
if audio.size == 0:
raise ValueError("empty audio")
if sample_rate != 16000:
audio = librosa.resample(audio, orig_sr=sample_rate, target_sr=16000)
pipeline_instance = self._get_pipeline(model_id=model_id, model_revision=model_revision)
model = getattr(pipeline_instance, "model", None)
if model is None:
raise RuntimeError("speaker verification model is not initialized")
device = getattr(pipeline_instance, "device_name", self._device())
with self._inference_lock:
with torch.no_grad():
embeddings = model(torch.as_tensor(audio[None, :]).to(device))
if isinstance(embeddings, torch.Tensor):
embedding = embeddings.detach().cpu().numpy()[0]
else:
embedding = np.asarray(embeddings, dtype=np.float32)[0]
return _normalize_embedding(embedding)
def extract_embedding_from_file(self, file_path: str) -> np.ndarray:
audio_data, sample_rate = librosa.load(file_path, sr=16000, mono=True)
return self.extract_embedding_from_audio(audio_data, int(sample_rate))
async def register_file(
self,
*,
name: str,
file_path: str,
user_id: Optional[str] = None,
) -> dict[str, Optional[str]]:
if not pg_speaker_db.is_connected:
raise RuntimeError("Speaker database is not connected")
loop = asyncio.get_running_loop()
embedding = await loop.run_in_executor(
None,
self.extract_embedding_from_file,
file_path,
)
speaker = await pg_speaker_db.save_speaker(name, embedding, user_id=user_id)
return {
"speaker_id": speaker.get("id"),
"name": speaker.get("name") or name,
"user_id": speaker.get("user_id"),
"speaker_model": "CampPlus",
}
async def identify_embedding(
self,
embedding: np.ndarray,
threshold: Optional[float] = None,
) -> dict[str, Optional[str]]:
if not pg_speaker_db.is_connected:
return {"speaker_id": None, "name": None, "user_id": None}
speaker = await pg_speaker_db.identify_speaker(
_normalize_embedding(embedding),
threshold if threshold is not None else settings.SV_THRESHOLD,
)
return {
"speaker_id": speaker.get("id"),
"name": speaker.get("name"),
"user_id": speaker.get("user_id"),
"speaker_model": "CampPlus",
}
async def identify_file(
self,
*,
file_path: str,
threshold: Optional[float] = None,
) -> dict[str, Optional[str]]:
loop = asyncio.get_running_loop()
embedding = await loop.run_in_executor(
None,
self.extract_embedding_from_file,
file_path,
)
return await self.identify_embedding(embedding, threshold=threshold)
async def apply_registered_speakers(
self,
result: Any,
threshold: Optional[float] = None,
) -> Any:
if not pg_speaker_db.is_connected:
return result
cache: dict[str, dict[str, Optional[str]]] = {}
for segment in getattr(result, "segments", []) or []:
embedding = getattr(segment, "speaker_embedding", None)
local_speaker_id = getattr(segment, "speaker_id", None)
if embedding is None or not local_speaker_id:
continue
if local_speaker_id not in cache:
matched_candidate = await self.identify_embedding(
np.asarray(embedding, dtype=np.float32),
threshold=threshold,
)
if matched_candidate.get("name"):
cache[local_speaker_id] = matched_candidate
else:
continue
matched = cache[local_speaker_id]
if matched.get("name"):
segment.speaker_id = matched.get("speaker_id") or segment.speaker_id
segment.speaker_name = matched.get("name")
segment.user_id = matched.get("user_id")
return result
async def list_speakers(self) -> list[dict[str, Optional[str]]]:
return await pg_speaker_db.list_speakers()
async def delete_speaker(self, speaker_id: str) -> bool:
if not speaker_id.isdigit():
return False
return await pg_speaker_db.delete_speaker(int(speaker_id))
@staticmethod
async def save_upload_to_temp(content: bytes, suffix: str = ".wav") -> str:
fd, path = tempfile.mkstemp(prefix="speaker_", suffix=suffix, dir=settings.TEMP_DIR)
try:
with os.fdopen(fd, "wb") as file_obj:
file_obj.write(content)
except Exception:
os.close(fd)
raise
return path
@staticmethod
def cleanup_file(path: Optional[str]) -> None:
if path and os.path.exists(path):
try:
os.remove(path)
except Exception as exc:
logger.warning("清理临时声纹文件失败 %s: %s", path, exc)
@staticmethod
def save_audio_array_to_temp(audio_data: np.ndarray, sample_rate: int = 16000) -> str:
fd, path = tempfile.mkstemp(prefix="speaker_array_", suffix=".wav", dir=settings.TEMP_DIR)
os.close(fd)
sf.write(path, np.asarray(audio_data, dtype=np.float32), sample_rate)
return path
_speaker_registry_service: Optional[SpeakerRegistryService] = None
_speaker_registry_lock = threading.Lock()
def get_speaker_registry_service() -> SpeakerRegistryService:
global _speaker_registry_service
if _speaker_registry_service is None:
with _speaker_registry_lock:
if _speaker_registry_service is None:
_speaker_registry_service = SpeakerRegistryService()
return _speaker_registry_service

View File

@ -0,0 +1,24 @@
# -*- coding: utf-8 -*-
"""
工具模块
包含通用工具函数和辅助功能
"""
from .common import generate_task_id, validate_text_input, parse_language_code
from .audio import save_audio_array, load_audio_file, generate_temp_audio_path, cleanup_temp_file
from .text_processing import apply_itn_to_text, normalize_asr_text
__all__ = [
# 通用工具函数
"generate_task_id",
"validate_text_input",
"parse_language_code",
# 音频工具函数
"save_audio_array",
"load_audio_file",
"generate_temp_audio_path",
"cleanup_temp_file",
# ITN(逆文本标准化)功能 - 基于itntext
"apply_itn_to_text",
"normalize_asr_text",
]

528
app/utils/audio.py 100644
View File

@ -0,0 +1,528 @@
# -*- coding: utf-8 -*-
"""
统一音频处理工具
ASR音频处理功能
"""
import os
import tempfile
import requests
import librosa
import soundfile as sf
import numpy as np
import subprocess
import logging
from dataclasses import dataclass
from typing import Tuple, Optional
from io import BytesIO
from urllib.parse import unquote, urlparse
from ..core.config import settings
from ..core.exceptions import (
InvalidParameterException,
InvalidMessageException,
DefaultServerErrorException,
)
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class NormalizedAudio:
path: str
timestamp_scale: float = 1.0
def download_audio_from_url(url: str, max_size: Optional[int] = None) -> bytes:
"""从 URL 或服务端本地路径读取音频文件
Args:
url: 音频文件 URL、本地路径或 file:// 路径
max_size: 最大文件大小限制
Returns:
音频文件的二进制数据
Raises:
InvalidParameterException: URL无效或下载失败
InvalidMessageException: 文件太大
"""
if not url:
raise InvalidParameterException("URL不能为空")
max_file_size = max_size or settings.MAX_AUDIO_SIZE
parsed = urlparse(url)
if parsed.scheme in {"", "file"}:
local_path = os.path.expanduser(unquote(parsed.path) if parsed.scheme == "file" else url)
if not os.path.isfile(local_path):
raise InvalidParameterException(f"本地音频文件不存在: {local_path}")
file_size = os.path.getsize(local_path)
if file_size > max_file_size:
max_size_mb = max_file_size // 1024 // 1024
raise InvalidMessageException(f"音频文件太大,最大支持{max_size_mb}MB")
with open(local_path, "rb") as file_obj:
return file_obj.read()
try:
response = requests.get(url, timeout=30, stream=True)
response.raise_for_status()
# 检查Content-Length头
content_length = response.headers.get("content-length")
if content_length and int(content_length) > max_file_size:
max_size_mb = max_file_size // 1024 // 1024
raise InvalidMessageException(f"音频文件太大,最大支持{max_size_mb}MB")
# 分块下载并检查大小
audio_data = BytesIO()
downloaded_size = 0
for chunk in response.iter_content(chunk_size=8192):
downloaded_size += len(chunk)
if downloaded_size > max_file_size:
max_size_mb = max_file_size // 1024 // 1024
raise InvalidMessageException(f"音频文件太大,最大支持{max_size_mb}MB")
audio_data.write(chunk)
return audio_data.getvalue()
except requests.RequestException as e:
raise InvalidParameterException(f"下载音频文件失败: {str(e)}")
def save_audio_to_temp_file(audio_data: bytes, suffix: str = ".wav") -> str:
"""保存音频数据到临时文件
Args:
audio_data: 音频二进制数据
suffix: 文件后缀
Returns:
临时文件路径
Raises:
AudioProcessingException: 保存失败
"""
try:
with tempfile.NamedTemporaryFile(
delete=False, suffix=suffix, dir=settings.TEMP_DIR
) as temp_file:
temp_file.write(audio_data)
return temp_file.name
except Exception as e:
raise DefaultServerErrorException(f"保存音频文件失败: {str(e)}")
def cleanup_temp_file(file_path: str) -> None:
"""清理临时文件
Args:
file_path: 文件路径
"""
try:
if file_path and os.path.exists(file_path):
os.remove(file_path)
except Exception:
# 静默忽略清理错误
pass
def load_audio_file(audio_path: str, target_sr: int = 16000) -> Tuple[np.ndarray, int]:
"""加载音频文件并转换为指定采样率
Args:
audio_path: 音频文件路径
target_sr: 目标采样率
Returns:
(audio_data, sample_rate): 音频数据和采样率
Raises:
AudioProcessingException: 加载失败
"""
try:
# 使用librosa加载音频
audio_data, sr = librosa.load(audio_path, sr=target_sr)
return audio_data, int(sr)
except Exception as e:
raise DefaultServerErrorException(f"加载音频文件失败: {str(e)}")
def get_audio_duration(audio_path: str) -> float:
"""获取音频文件时长
Args:
audio_path: 音频文件路径
Returns:
音频时长(秒)
Raises:
AudioProcessingException: 获取时长失败
"""
try:
# Load audio and get duration
y, sr = librosa.load(audio_path, sr=None)
duration = librosa.get_duration(y=y, sr=sr)
return duration
except Exception as e:
raise DefaultServerErrorException(f"获取音频时长失败: {str(e)}")
def get_container_duration(audio_path: str) -> Optional[float]:
"""通过 ffprobe 获取音频容器的 metadata 时长
对于 m4a/AAC 等压缩格式,容器记录的时长可能与实际解码样本数不一致
(常见于 m3u8/ts 分片合并的音频)。返回 None 表示获取失败。
"""
try:
result = subprocess.run(
["ffprobe", "-v", "error", "-show_entries", "format=duration",
"-of", "default=noprint_wrappers=1:nokey=1", audio_path],
capture_output=True, text=True, timeout=10,
)
if result.returncode == 0 and result.stdout.strip():
return float(result.stdout.strip())
except Exception as e:
logger.debug(f"ffprobe 获取容器时长失败: {e}")
return None
def get_timestamp_scale(original_audio_path: str, decoded_duration: float) -> float:
"""计算时间戳缩放系数
对比容器 metadata 时长与解码后实际时长,返回缩放系数。
用于修正 m4a/AAC 等格式中容器时长与解码时长不一致的问题。
Args:
original_audio_path: 原始音频文件路径(转换前)
decoded_duration: 解码后的实际音频时长(秒)
Returns:
缩放系数(容器时长 / 解码时长),无差异时返回 1.0
"""
container_duration = get_container_duration(original_audio_path)
if container_duration is None or decoded_duration <= 0:
return 1.0
scale = container_duration / decoded_duration
if abs(scale - 1.0) < 0.001:
# 差异 < 0.1%,忽略
return 1.0
logger.info(
f"检测到容器/解码时长不一致: container={container_duration:.3f}s, "
f"decoded={decoded_duration:.3f}s, scale={scale:.6f}"
)
return scale
def resample_audio_array(
audio_array: np.ndarray,
original_sr: int,
target_sr: int,
) -> np.ndarray:
"""重采样音频数组
Args:
audio_array: 原始音频数据
original_sr: 原始采样率
target_sr: 目标采样率
Returns:
重采样后的音频数据
"""
if original_sr == target_sr:
return audio_array
try:
# 确保是1D数组用于librosa重采样
if audio_array.ndim > 1:
# 如果是多声道,取第一个声道
if audio_array.shape[0] > audio_array.shape[1]:
audio_1d = audio_array[0, :]
else:
audio_1d = (
audio_array[:, 0]
if audio_array.shape[1] > 1
else audio_array.flatten()
)
else:
audio_1d = audio_array
# 使用librosa进行重采样
resampled = librosa.resample(audio_1d, orig_sr=original_sr, target_sr=target_sr)
logger.info(f"音频重采样: {original_sr}Hz -> {target_sr}Hz")
return resampled
except Exception as e:
logger.warning(f"音频重采样失败: {str(e)},使用原始音频")
return audio_array
def adjust_audio_volume(audio_array: np.ndarray, volume: int) -> np.ndarray:
"""调节音频音量
Args:
audio_array: 音频数据数组
volume: 音量值,范围0~100,50为原始音量
Returns:
调节后的音频数据
"""
if int(volume) == 50:
return audio_array
if volume < 0 or volume > 100:
logger.warning(f"音量值{volume}超出范围[0,100],使用默认值50")
volume = 50
# 将音量值转换为倍数 (0-100 -> 0-2.0)
volume_factor = volume / 50.0
# 应用音量调节
adjusted_audio = audio_array * volume_factor
# 防止削波,如果音量过大导致超过范围,进行归一化
max_val = np.max(np.abs(adjusted_audio))
if max_val > 1.0:
adjusted_audio = adjusted_audio / max_val
logger.info(f"音量调节后进行归一化,最大值: {max_val:.3f}")
logger.info(f"音频音量已调节: {volume}/100 (倍数: {volume_factor:.2f})")
return adjusted_audio
def save_audio_array(
audio_array: np.ndarray,
output_path: str,
sample_rate: int = 22050,
format: str = "wav",
original_sr: Optional[int] = None,
volume: int = 50,
) -> str:
"""保存音频数组到文件
Args:
audio_array: 音频数据数组
output_path: 输出文件路径
sample_rate: 目标采样率
format: 音频格式
original_sr: 原始采样率(用于重采样)
volume: 音量值,范围0~100,默认50
Returns:
保存的文件路径
Raises:
AudioProcessingException: 保存失败
"""
try:
# 如果指定了原始采样率且与目标采样率不同,进行重采样
if original_sr and original_sr != sample_rate:
audio_array = resample_audio_array(audio_array, original_sr, sample_rate)
# 调节音频音量
audio_array = adjust_audio_volume(audio_array, volume)
# 确保音频数据是float32格式
if audio_array.dtype != np.float32:
audio_array = audio_array.astype(np.float32)
# 确保音频数据在正确的范围内
if np.max(np.abs(audio_array)) > 1.0:
audio_array = audio_array / np.max(np.abs(audio_array))
# 确保是2D张量 (channels, samples)
if audio_array.ndim == 1:
audio_array = audio_array[np.newaxis, :] # 添加通道维度
elif audio_array.ndim > 2:
audio_array = audio_array.squeeze()
if audio_array.ndim == 1:
audio_array = audio_array[np.newaxis, :]
# 根据格式选择保存方法
if format.lower() == "wav":
sf.write(output_path, audio_array.T, sample_rate, format="WAV")
else:
# 使用soundfile保存其他格式
# 确保音频数据是单声道
if audio_array.shape[0] > 1:
audio_array = np.mean(audio_array, axis=0)
sf.write(output_path, audio_array.T, sample_rate, format=format.upper())
return output_path
except Exception as e:
raise DefaultServerErrorException(f"保存音频文件失败: {str(e)}")
def convert_audio_to_wav(
input_path: str, output_path: Optional[str] = None, target_sr: int = 16000
) -> str:
"""转换音频文件为WAV格式
Args:
input_path: 输入文件路径
output_path: 输出文件路径(可选)
target_sr: 目标采样率,默认16000Hz
Returns:
转换后的文件路径
Raises:
AudioProcessingException: 转换失败
"""
if not output_path:
output_path = input_path.rsplit(".", 1)[0] + ".wav"
try:
# 使用librosa加载并重采样
audio_data, _ = librosa.load(input_path, sr=target_sr)
sf.write(output_path, audio_data, target_sr, format="WAV")
return output_path
except Exception as e:
# 尝试使用ffmpeg转换
try:
subprocess.run(
[
"ffmpeg",
"-f", "s16le",
"-ar", str(target_sr),
"-ac", "1",
"-i", input_path,
"-acodec", "pcm_s16le",
output_path,
"-y",
],
check=True,
capture_output=True,
)
return output_path
except (subprocess.CalledProcessError, FileNotFoundError):
raise DefaultServerErrorException(f"音频格式转换失败: {str(e)}")
def normalize_audio_for_asr(audio_path: str, target_sr: int = 16000) -> NormalizedAudio:
"""Normalize audio and return explicit timestamp metadata.
Args:
audio_path: 输入音频文件路径
target_sr: 目标采样率,默认16000Hz
Returns:
Normalized audio path and timestamp scale metadata.
"""
try:
# 检查文件扩展名
file_ext = os.path.splitext(audio_path)[1].lower()
# 如果已经是WAV格式且采样率正确,直接返回
if file_ext == ".wav":
# 检查采样率
_, sr = librosa.load(audio_path, sr=None)
if sr == target_sr:
return NormalizedAudio(path=audio_path)
# 转换为标准WAV格式
normalized_path = convert_audio_to_wav(audio_path, target_sr=target_sr)
logger.debug(f"音频文件已标准化: {audio_path} -> {normalized_path}")
timestamp_scale = 1.0
if normalized_path != audio_path:
decoded_duration = get_audio_duration(normalized_path)
timestamp_scale = get_timestamp_scale(audio_path, decoded_duration)
return NormalizedAudio(path=normalized_path, timestamp_scale=timestamp_scale)
except Exception as e:
raise DefaultServerErrorException(f"音频标准化失败: {str(e)}")
def generate_temp_audio_path(prefix: str = "audio", suffix: str = ".wav") -> str:
"""生成临时音频文件路径
Args:
prefix: 文件名前缀
suffix: 文件后缀
Returns:
临时文件路径
"""
import time
timestamp = int(time.time())
filename = f"{prefix}_{timestamp}_{os.getpid()}{suffix}"
return os.path.join(settings.TEMP_DIR, filename)
def detect_audio_format_from_bytes(data: bytes) -> str:
"""通过文件头(magic bytes)检测音频格式
Args:
data: 音频文件的前几个字节
Returns:
文件后缀(包含点号)
"""
if len(data) < 12:
return ".wav"
# 检查常见音频格式的文件头
if data[:4] == b"RIFF" and data[8:12] == b"WAVE":
return ".wav"
elif data[:3] == b"ID3" or (data[0:2] == b"\xff\xfb") or (data[0:2] == b"\xff\xfa"):
return ".mp3"
elif data[:4] == b"fLaC":
return ".flac"
elif data[:4] == b"OggS":
return ".ogg"
elif data[4:8] == b"ftyp":
# M4A/AAC/MP4/MOV 容器
return ".mp4"
elif data[:4] == b"\x1aE\xdf\xa3":
# WebM/MKV
return ".webm"
# 默认为 wav,librosa 会自动处理
return ".wav"
def get_audio_file_suffix(
audio_address: Optional[str] = None, audio_data: Optional[bytes] = None
) -> str:
"""自动识别音频文件后缀
Args:
audio_address: 音频文件URL(可选)
audio_data: 音频二进制数据(可选,用于检测文件头)
Returns:
文件后缀(包含点号)
"""
if audio_address:
# 从URL中提取扩展名
parsed = urlparse(audio_address)
path = unquote(parsed.path)
# 获取扩展名
ext = os.path.splitext(path)[1].lower()
if ext and ext in [
".wav", ".mp3", ".flac", ".ogg", ".m4a", ".aac", ".pcm", ".webm",
".mp4", ".mpeg", ".mpga", ".mov", ".mkv", ".avi",
]:
return ext
# 无法识别扩展名,默认为 .wav
return ".wav"
elif audio_data:
# 通过文件头检测格式
return detect_audio_format_from_bytes(audio_data[:12])
else:
# 默认为 .wav
return ".wav"

View File

@ -0,0 +1,64 @@
# -*- coding: utf-8 -*-
"""
音频过滤工具 - 用于流式ASR的近场/远场声音检测
"""
import numpy as np
import logging
from typing import Tuple, Dict
logger = logging.getLogger(__name__)
def calculate_rms_energy(audio_array: np.ndarray) -> float:
"""计算音频RMS能量
Args:
audio_array: float32音频数组,范围-1.0到1.0
Returns:
RMS能量值
"""
if len(audio_array) == 0:
return 0.0
return float(np.sqrt(np.mean(audio_array ** 2)))
def is_nearfield_voice(
audio_array: np.ndarray,
sample_rate: int = 16000, # noqa: ARG001
rms_threshold: float = 0.01,
enable_filter: bool = True,
) -> Tuple[bool, Dict]:
"""判断是否为近场有效声音(仅基于RMS能量)
Args:
audio_array: float32音频数组,范围-1.0到1.0
sample_rate: 采样率(保留用于兼容性)
rms_threshold: RMS能量阈值
enable_filter: 是否启用过滤(开关)
Returns:
(is_nearfield, metrics): 是否近场声音 + 检测指标详情
"""
if not enable_filter:
return True, {'enabled': False}
if len(audio_array) == 0:
return False, {'error': 'empty_array'}
# 计算RMS能量
rms_energy = calculate_rms_energy(audio_array)
# 仅使用RMS能量判断
is_nearfield = rms_energy >= rms_threshold
metrics = {
'rms_energy': round(rms_energy, 6),
'is_nearfield': is_nearfield,
'thresholds': {
'rms': rms_threshold,
}
}
return is_nearfield, metrics

View File

@ -0,0 +1,439 @@
# -*- coding: utf-8 -*-
"""
音频分割模块
基于 VAD 的智能音频分割,支持长音频分段识别
"""
import logging
import numpy as np
import librosa
import soundfile as sf
import tempfile
import os
import time
from typing import List, Tuple, Optional
from dataclasses import dataclass
from ..core.config import settings
from ..core.exceptions import DefaultServerErrorException
logger = logging.getLogger(__name__)
def _log_audio_split_timing(stage: str, duration_ms: float, **extra) -> None:
payload = {
"event": "audio_split_timing",
"stage": stage,
"duration_ms": round(duration_ms, 2),
}
payload.update(extra)
logger.info("音频分割阶段耗时", extra=payload)
@dataclass
class AudioSegment:
"""音频片段信息"""
start_ms: int # 开始时间(毫秒)
end_ms: int # 结束时间(毫秒)
audio_data: Optional[np.ndarray] = None # 音频数据
temp_file: Optional[str] = None # 临时文件路径
speaker_id: Optional[str] = None # 说话人ID(多说话人模式)
@property
def start_sec(self) -> float:
"""开始时间(秒)"""
return self.start_ms / 1000.0
@property
def end_sec(self) -> float:
"""结束时间(秒)"""
return self.end_ms / 1000.0
@property
def duration_ms(self) -> int:
"""时长(毫秒)"""
return self.end_ms - self.start_ms
@property
def duration_sec(self) -> float:
"""时长(秒)"""
return self.duration_ms / 1000.0
class AudioSplitter:
"""音频分割器
使用 VAD 模型检测语音边界,智能分割长音频
"""
# 默认配置
DEFAULT_MIN_SEGMENT_SEC = 1.0 # 每段最小时长(秒)
DEFAULT_SAMPLE_RATE = 16000 # 默认采样率
def __init__(
self,
min_segment_sec: float = DEFAULT_MIN_SEGMENT_SEC,
device: str = "auto",
):
"""初始化音频分割器
Args:
min_segment_sec: 每段最小时长(秒)
device: 计算设备("cuda", "cpu", "auto")
"""
split_trigger_sec = settings.MAX_SEGMENT_SEC
self.split_trigger_sec = split_trigger_sec
self.min_segment_sec = min_segment_sec
self.split_trigger_ms = int(split_trigger_sec * 1000)
self.min_segment_ms = int(min_segment_sec * 1000)
self.device = device
def get_vad_segments(
self, audio_path: str
) -> List[Tuple[int, int]]:
"""使用 VAD 模型获取语音段
Args:
audio_path: 音频文件路径
Returns:
语音段列表,每个元素为 (start_ms, end_ms)
"""
try:
from ..services.asr.engines import get_global_vad_model
logger.info("开始 VAD 语音段检测...")
vad_model = get_global_vad_model(self.device)
if vad_model is None:
raise DefaultServerErrorException("VAD 模型未加载")
# 调用 VAD 模型
vad_started = time.perf_counter()
result = vad_model.generate(input=audio_path, cache={})
vad_duration_ms = (time.perf_counter() - vad_started) * 1000
if not result or len(result) == 0:
_log_audio_split_timing(
"vad_generate",
vad_duration_ms,
audio_path=audio_path,
vad_segment_count=0,
)
logger.warning("VAD 未检测到语音段")
return []
# 解析 VAD 结果
# FunASR VAD 返回格式: [[start_ms, end_ms], [start_ms, end_ms], ...]
vad_segments = result[0].get("value", [])
if not vad_segments:
_log_audio_split_timing(
"vad_generate",
vad_duration_ms,
audio_path=audio_path,
vad_segment_count=0,
)
logger.warning("VAD 结果为空")
return []
_log_audio_split_timing(
"vad_generate",
vad_duration_ms,
audio_path=audio_path,
vad_segment_count=len(vad_segments),
)
logger.info(f"VAD 检测到 {len(vad_segments)} 个语音段")
logger.info(
"开始按 VAD 边界重分段 "
f"(split_trigger={self.split_trigger_sec}s, min_segment={self.min_segment_sec}s)..."
)
return [(int(seg[0]), int(seg[1])) for seg in vad_segments]
except Exception as e:
logger.error(f"VAD 检测失败: {e}")
raise DefaultServerErrorException(f"VAD 检测失败: {str(e)}")
def merge_segments_greedy(
self, vad_segments: List[Tuple[int, int]], total_duration_ms: int
) -> List[Tuple[int, int]]:
"""按 VAD 结果重分段
策略:
1. 默认保留 VAD 原始边界,避免将整段连续语音合并成超长片段
2. 仅对短片段(< min_segment_ms)做邻段合并
3. 对重叠片段进行边界修正,避免重复音频
Args:
vad_segments: VAD 检测到的语音段列表 [(start_ms, end_ms), ...]
total_duration_ms: 音频总时长(毫秒)
Returns:
合并后的段列表 [(start_ms, end_ms), ...]
"""
if not vad_segments:
# 没有 VAD 段,返回整个音频(按最大时长切分)
return self._split_by_fixed_duration(total_duration_ms)
# 按时间排序并修正边界(防止越界、重叠)
sorted_vad = sorted(vad_segments, key=lambda x: x[0])
normalized: List[Tuple[int, int]] = []
for raw_start, raw_end in sorted_vad:
start_ms = max(0, int(raw_start))
end_ms = min(total_duration_ms, int(raw_end))
if end_ms <= start_ms:
continue
if not normalized:
normalized.append((start_ms, end_ms))
continue
last_end = normalized[-1][1]
# 有重叠时,优先保持边界,避免与上一段重复采样
if start_ms < last_end:
start_ms = last_end
if end_ms > start_ms:
normalized.append((start_ms, end_ms))
if not normalized:
return self._split_by_fixed_duration(total_duration_ms)
merged = list(normalized)
# 只处理短片段:与相邻片段合并(不基于静音间隙)
idx = 0
while idx < len(merged):
start_ms, end_ms = merged[idx]
duration = end_ms - start_ms
if duration >= self.min_segment_ms or len(merged) == 1:
idx += 1
continue
if idx == 0:
# 首段过短:并入后段
next_end = merged[idx + 1][1]
merged[idx + 1] = (start_ms, next_end)
del merged[idx]
continue
if idx == len(merged) - 1:
# 尾段过短:并入前段
prev_start, _ = merged[idx - 1]
merged[idx - 1] = (prev_start, end_ms)
del merged[idx]
idx = max(0, idx - 1)
continue
# 中间短段:优先并入时长更短的一侧,避免单段过长
prev_start = merged[idx - 1][0]
next_end = merged[idx + 1][1]
merged_with_prev_duration = end_ms - prev_start
merged_with_next_duration = next_end - start_ms
if merged_with_prev_duration <= merged_with_next_duration:
merged[idx - 1] = (prev_start, end_ms)
del merged[idx]
idx = max(0, idx - 1)
else:
merged[idx + 1] = (start_ms, next_end)
del merged[idx]
return merged
def _split_by_fixed_duration(self, total_duration_ms: int) -> List[Tuple[int, int]]:
"""按固定时长切分(无 VAD 时的 fallback)
Args:
total_duration_ms: 音频总时长(毫秒)
Returns:
切分后的段列表
"""
segments = []
current = 0
while current < total_duration_ms:
end = min(current + self.split_trigger_ms, total_duration_ms)
if end - current >= self.min_segment_ms:
segments.append((current, end))
current = end
return segments
def split_audio_file(
self,
audio_path: str,
output_dir: Optional[str] = None,
) -> List[AudioSegment]:
"""分割音频文件
Args:
audio_path: 音频文件路径
output_dir: 输出目录(可选,默认使用临时目录)
Returns:
音频片段列表
"""
try:
total_started = time.perf_counter()
# 加载音频
load_started = time.perf_counter()
audio_data, sr = librosa.load(audio_path, sr=self.DEFAULT_SAMPLE_RATE)
load_audio_ms = (time.perf_counter() - load_started) * 1000
total_duration_ms = int(len(audio_data) / sr * 1000)
audio_duration_sec = total_duration_ms / 1000
logger.info(f"音频总时长: {audio_duration_sec:.2f}秒")
_log_audio_split_timing(
"load_audio",
load_audio_ms,
audio_path=audio_path,
audio_duration_sec=round(audio_duration_sec, 2),
sample_rate=sr,
)
# 检查是否需要分割
if total_duration_ms <= self.split_trigger_ms:
logger.info("音频时长在限制内,无需分割")
total_duration_ms_for_log = (time.perf_counter() - total_started) * 1000
_log_audio_split_timing(
"split_total",
total_duration_ms_for_log,
audio_path=audio_path,
audio_duration_sec=round(audio_duration_sec, 2),
output_segment_count=1,
need_split=False,
load_audio_ms=round(load_audio_ms, 2),
vad_ms=0,
merge_ms=0,
write_segments_ms=0,
)
return [
AudioSegment(
start_ms=0,
end_ms=total_duration_ms,
audio_data=audio_data,
temp_file=audio_path,
)
]
# 获取 VAD 段
vad_started = time.perf_counter()
vad_segments = self.get_vad_segments(audio_path)
vad_ms = (time.perf_counter() - vad_started) * 1000
# 贪婪合并
merge_started = time.perf_counter()
merged_segments = self.merge_segments_greedy(vad_segments, total_duration_ms)
merge_ms = (time.perf_counter() - merge_started) * 1000
logger.info(f"重分段完成: 原始VAD={len(vad_segments)}, 输出={len(merged_segments)}")
_log_audio_split_timing(
"merge_segments",
merge_ms,
audio_path=audio_path,
vad_segment_count=len(vad_segments),
output_segment_count=len(merged_segments),
audio_duration_sec=round(audio_duration_sec, 2),
)
# 切分音频并保存到临时文件
logger.info("开始切分音频并保存临时文件...")
output_dir = output_dir or settings.TEMP_DIR
os.makedirs(output_dir, exist_ok=True)
audio_segments = []
write_started = time.perf_counter()
for idx, (start_ms, end_ms) in enumerate(merged_segments):
# 计算采样点范围
start_sample = int(start_ms / 1000 * sr)
end_sample = int(end_ms / 1000 * sr)
# 提取音频片段
segment_data = audio_data[start_sample:end_sample]
# 保存到临时文件
temp_file = tempfile.NamedTemporaryFile(
delete=False,
suffix=".wav",
dir=output_dir,
prefix=f"segment_{idx:03d}_",
)
temp_path = temp_file.name
temp_file.close()
sf.write(temp_path, segment_data, sr)
segment = AudioSegment(
start_ms=start_ms,
end_ms=end_ms,
audio_data=segment_data,
temp_file=temp_path,
)
audio_segments.append(segment)
logger.debug(
f"分段 {idx + 1}/{len(merged_segments)}: "
f"{start_ms / 1000:.2f}s - {end_ms / 1000:.2f}s "
f"(时长: {segment.duration_sec:.2f}s)"
)
write_segments_ms = (time.perf_counter() - write_started) * 1000
logger.info(f"音频切分完成,共 {len(audio_segments)} 个分段")
_log_audio_split_timing(
"write_segments",
write_segments_ms,
audio_path=audio_path,
output_dir=output_dir,
output_segment_count=len(audio_segments),
audio_duration_sec=round(audio_duration_sec, 2),
)
total_ms = (time.perf_counter() - total_started) * 1000
_log_audio_split_timing(
"split_total",
total_ms,
audio_path=audio_path,
audio_duration_sec=round(audio_duration_sec, 2),
output_segment_count=len(audio_segments),
need_split=True,
load_audio_ms=round(load_audio_ms, 2),
vad_ms=round(vad_ms, 2),
merge_ms=round(merge_ms, 2),
write_segments_ms=round(write_segments_ms, 2),
)
return audio_segments
except Exception as e:
logger.error(f"音频分割失败: {e}")
raise DefaultServerErrorException(f"音频分割失败: {str(e)}")
@staticmethod
def cleanup_segments(segments: List[AudioSegment]) -> None:
"""清理临时文件
Args:
segments: 音频片段列表
"""
for segment in segments:
if segment.temp_file and os.path.exists(segment.temp_file):
try:
os.remove(segment.temp_file)
except Exception as e:
logger.warning(f"清理临时文件失败: {segment.temp_file}, {e}")
def split_long_audio(
audio_path: str,
device: str = "auto",
) -> List[AudioSegment]:
"""分割长音频的便捷函数
Args:
audio_path: 音频文件路径
device: 计算设备
Returns:
音频片段列表
"""
splitter = AudioSplitter(device=device)
return splitter.split_audio_file(audio_path)

View File

@ -0,0 +1,28 @@
# -*- coding: utf-8 -*-
"""Structured startup events for the optional terminal dashboard."""
from __future__ import annotations
import json
import os
import sys
from typing import Any
_BOOT_EVENT_PREFIX = "__FUNASR_BOOT__"
def boot_events_enabled() -> bool:
return (os.getenv("FUNASR_BOOT_EVENTS") or "").strip() == "1"
def emit_boot_event(event: str, **payload: Any) -> None:
if not boot_events_enabled():
return
data = {"event": event, **payload}
sys.stdout.write(_BOOT_EVENT_PREFIX + json.dumps(data, ensure_ascii=False) + "\n")
sys.stdout.flush()
def get_boot_event_prefix() -> str:
return _BOOT_EVENT_PREFIX

View File

@ -0,0 +1,90 @@
# -*- coding: utf-8 -*-
"""
通用工具函数
包含任务ID生成、参数验证等通用功能
"""
import uuid
import hashlib
import time
import re
from typing import Optional
def generate_task_id(prefix: str = "") -> str:
"""生成唯一的任务ID
Args:
prefix: 任务ID前缀
Returns:
生成的任务ID
"""
timestamp = str(int(time.time() * 1000))
random_id = str(uuid.uuid4()).replace("-", "")
combined = timestamp + random_id
# 使用MD5哈希生成32位字符串
task_id = hashlib.md5(combined.encode()).hexdigest()
if prefix:
return f"{prefix}_{task_id}"
return task_id
def validate_text_input(text: str, max_length: int = 10000) -> tuple[bool, str]:
"""验证输入文本
Args:
text: 待验证的文本
max_length: 最大长度限制
Returns:
(is_valid, message): 验证结果和消息
"""
if not text or not text.strip():
return False, "文本内容不能为空"
text = text.strip()
if len(text) > max_length:
return False, f"文本长度超过限制,最大支持{max_length}个字符"
# 检查是否包含有效字符
if not re.search(r"[\u4e00-\u9fff\w\s]", text):
return False, "文本内容无效,请输入有效的中文、英文或数字"
return True, "验证通过"
def parse_language_code(lang_code: Optional[str]) -> str:
"""解析语言代码
Args:
lang_code: 语言代码(如 zh, zh-cn, en, ja等)
Returns:
标准化的语言代码
"""
if not lang_code:
return "zh" # 默认中文
lang_code = lang_code.lower().strip()
# 语言代码映射
lang_mapping = {
"zh": "zh",
"zh-cn": "zh",
"zh-tw": "zh",
"zh-hk": "zh",
"en": "en",
"en-us": "en",
"en-gb": "en",
"ja": "jp",
"jp": "jp",
"ko": "kr",
"kr": "kr",
"yue": "yue", # 粤语
}
return lang_mapping.get(lang_code, "zh")

View File

@ -0,0 +1,296 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
模型预下载脚本
统一从 ModelScope 预下载所有运行所需模型
"""
import argparse
import json
from pathlib import Path
from typing import Optional
from modelscope.hub.snapshot_download import snapshot_download as ms_snapshot_download
from app.core.config import settings
from app.services.asr.model_capabilities import (
get_all_qwen_modelscope_assets,
get_camplusplus_replacement_paths,
get_download_modelscope_assets,
)
def _get_qwen_modelscope_assets():
"""Return all declared ModelScope Qwen assets for offline deployment."""
assets = get_all_qwen_modelscope_assets()
if not assets:
print("当前部署计划未启用 Qwen3-ASR,跳过 Qwen 模型下载")
return []
print("离线部署模式:下载全部已声明 Qwen 模型(含 forced aligner)")
return assets
def _get_cache_path(model_id: str, source: str = "modelscope") -> Path:
"""获取模型缓存路径"""
_ = source
return Path(settings.MODELSCOPE_PATH) / model_id
def check_model_exists(model_id: str, source: str = "modelscope") -> tuple[bool, str]:
"""检查模型是否已存在于本地缓存"""
try:
model_path = _get_cache_path(model_id, source)
if model_path.exists() and model_path.is_dir():
if any(model_path.iterdir()):
return True, str(model_path)
except Exception:
pass
return False, ""
def check_all_models() -> list[tuple[str, str, str, Optional[str]]]:
"""检查所有模型是否存在
Returns:
缺失的模型列表,每个元素为 (model_id, description, source, revision)
"""
missing = []
ms_assets = get_download_modelscope_assets()
qwen_assets = _get_qwen_modelscope_assets()
for asset in ms_assets:
exists, _ = check_model_exists(asset.model_id, source="modelscope")
if not exists:
missing.append((asset.model_id, asset.description, "modelscope", asset.revision))
for asset in qwen_assets:
exists, _ = check_model_exists(asset.model_id, source="modelscope")
if not exists:
missing.append((asset.model_id, asset.description, "modelscope", asset.revision))
return missing
def fix_camplusplus_config() -> bool:
"""修复 CAM++ 配置文件,将模型ID替换为本地路径(用于离线环境)
修复 issue #15: 离线环境下 CAM++ 模型会尝试从 modelscope.cn 获取依赖模型配置
Returns:
是否修复成功
"""
try:
cache_dir = Path(settings.MODELSCOPE_PATH)
config_file = cache_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json"
if not config_file.exists():
return False
# 读取配置文件
with open(config_file, 'r', encoding='utf-8') as f:
config = json.load(f)
# 需要替换的模型ID -> 本地路径映射
replacements = get_camplusplus_replacement_paths(str(cache_dir))
# 检查是否需要修改
modified = False
if "model" in config:
for key in ["speaker_model", "change_locator", "vad_model"]:
if key in config["model"]:
old_value = config["model"][key]
if old_value in replacements:
new_value = replacements[old_value]
# 检查本地路径是否存在
if Path(new_value).exists():
config["model"][key] = new_value
modified = True
# 写回配置文件
if modified:
with open(config_file, 'w', encoding='utf-8') as f:
json.dump(config, f, indent=4, ensure_ascii=False)
return True
return False
except Exception as e:
print(f"⚠️ 修复 CAM++ 配置文件失败: {e}")
return False
def download_models(
auto_mode: bool = False,
export_dir: Optional[str] = None,
) -> bool:
"""下载所有需要的模型
Args:
auto_mode: 如果为True,表示自动模式(从start.py调用),会简化输出
export_dir: 如果指定,将下载的模型导出到该目录(用于离线部署)
Returns:
是否全部下载成功
"""
import shutil
# 检查缺失的模型
missing = check_all_models()
ms_assets = get_download_modelscope_assets()
qwen_assets = _get_qwen_modelscope_assets()
export_path = Path(export_dir) if export_dir else None
if not missing:
if not auto_mode:
print("✅ 所有模型已存在,无需下载")
if not export_path:
return True
ms_cache_dir = Path(settings.MODELSCOPE_CACHE)
if auto_mode:
print(f"📦 检测到 {len(missing)} 个模型需要下载...")
else:
print("=" * 60)
print("Qwen3-ASR 模型预下载")
print("=" * 60)
print(f"ModelScope 缓存: {ms_cache_dir}")
print(f"待下载模型: {len(missing)} 个")
print("=" * 60)
failed = []
downloaded = []
# 下载 ModelScope 模型 (Paraformer)
ms_missing = [(mid, desc, rev) for mid, desc, src, rev in missing if src == "modelscope"]
if ms_missing:
if not auto_mode:
print("\n📦 开始下载 ModelScope 模型 (Paraformer)...")
print("-" * 60)
for i, (model_id, desc, revision) in enumerate(ms_missing, 1):
if not auto_mode:
print(f"\n[{i}/{len(ms_missing)}] {desc}")
print(f" 模型ID: {model_id}")
if revision:
print(f" 版本: {revision}")
print(f" 📥 开始下载...", end="")
try:
local_dir = _get_cache_path(model_id, "modelscope")
local_dir.parent.mkdir(parents=True, exist_ok=True)
# 传递版本参数,如果指定了版本
if revision:
path = ms_snapshot_download(
model_id,
revision=revision,
cache_dir=str(ms_cache_dir),
local_dir=str(local_dir),
)
else:
path = ms_snapshot_download(
model_id,
cache_dir=str(ms_cache_dir),
local_dir=str(local_dir),
)
if not auto_mode:
print(f" ✅ 完成: {path}")
downloaded.append((model_id, "modelscope", path))
except Exception as e:
if not auto_mode:
print(f" ❌ 失败: {e}")
failed.append((model_id, str(e)))
# 修复 CAM++ 配置文件(用于离线环境)
if not auto_mode:
print("\n🔧 修复 CAM++ 配置文件...")
if fix_camplusplus_config():
if not auto_mode:
print(" ✅ CAM++ 配置已修复(离线环境可用)")
else:
if not auto_mode:
print(" ℹ️ 无需修复或配置文件不存在")
# 导出模式:复制模型到扁平化的 models/ 根目录
if export_path and not failed:
if not auto_mode:
print(f"\n📦 导出模型到: {export_path}")
# 收集所有需要导出的模型
all_models = []
for asset in ms_assets:
all_models.append((asset.model_id, "modelscope"))
for asset in qwen_assets:
all_models.append((asset.model_id, "modelscope"))
exported = 0
for model_entry in all_models:
if len(model_entry) == 3:
model_id, source, actual_model_id = model_entry
else:
model_id, source = model_entry
actual_model_id = model_id
cache_path = _get_cache_path(actual_model_id, source)
if cache_path.exists():
rel_path = cache_path.relative_to(Path(settings.MODELSCOPE_PATH))
target_dir = export_path / rel_path
target_dir.parent.mkdir(parents=True, exist_ok=True)
if not auto_mode:
print(f" 📂 {model_id}", end="")
try:
shutil.copytree(cache_path, target_dir, dirs_exist_ok=True)
exported += 1
if not auto_mode:
print(" ✅")
except Exception as e:
if not auto_mode:
print(f" ❌ {e}")
if not auto_mode:
print(f"\n✅ 已导出 {exported} 个模型到 models/")
if not auto_mode:
print("\n" + "=" * 60)
print("📊 下载统计:")
print(f" ✅ 已下载: {len(downloaded)} 个")
print(f" ❌ 失败: {len(failed)} 个")
print("=" * 60)
if failed:
print(f"\n失败的模型:")
for model_id, err in failed:
print(f" - {model_id}: {err}")
return False
else:
print("\n✅ 所有模型准备就绪!")
print("=" * 60)
return len(failed) == 0
def main() -> int:
"""CLI entrypoint for model download and export."""
parser = argparse.ArgumentParser(description="Download or export Qwen3-ASR models")
parser.add_argument(
"--export-dir",
default=None,
help="Optional export directory for offline deployment packaging",
)
parser.add_argument(
"--auto-mode",
action="store_true",
help="Reduce output for startup/bootstrap usage",
)
args = parser.parse_args()
success = download_models(
auto_mode=args.auto_mode,
export_dir=args.export_dir,
)
return 0 if success else 1
if __name__ == "__main__":
raise SystemExit(main())

View File

@ -0,0 +1,520 @@
# -*- coding: utf-8 -*-
"""
模型预加载工具
在应用启动时预加载所有需要的模型,避免首次请求时的延迟
"""
import logging
import os
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any
try:
from rich.console import Console
except ImportError:
Console = None
from .boot_events import emit_boot_event
logger = logging.getLogger(__name__)
_PRELOAD_QUIET_LOGGERS = (
"root",
"vllm",
"app.infrastructure.model_utils",
"app.services.asr.engines.global_models",
"app.services.asr.qwen3_engine",
"app.utils.speaker_diarizer",
)
class _ProgressNoiseFilter(logging.Filter):
def filter(self, record: logging.LogRecord) -> bool:
if record.levelno >= logging.WARNING:
return True
return not any(
record.name == prefix or record.name.startswith(f"{prefix}.")
for prefix in _PRELOAD_QUIET_LOGGERS
)
class _StartupProgress:
def __init__(self, title: str, total: int):
self._title = title
self._total = max(total, 1)
self._enabled = bool(
Console is not None
and sys.stderr.isatty()
)
self._console: Any = None
self._filter = _ProgressNoiseFilter()
self._handlers: list[logging.Handler] = []
self._current_step = 1
self._last_description: str | None = None
def __enter__(self) -> "_StartupProgress":
emit_boot_event(
"phase_start",
phase=self._title,
total=self._total,
message=self._title,
)
if not self._enabled or Console is None:
return self
self._console = Console(stderr=True)
root_logger = logging.getLogger()
self._handlers = list(root_logger.handlers)
for handler in self._handlers:
handler.addFilter(self._filter)
return self
def __exit__(self, exc_type, exc, tb) -> None:
for handler in self._handlers:
handler.removeFilter(self._filter)
self._handlers.clear()
def update(self, description: str) -> None:
emit_boot_event(
"step_start",
phase=self._title,
step=self._current_step,
total=self._total,
message=description,
)
if self._console is None:
return
if description == self._last_description:
return
self._last_description = description
self._console.print(
f"[bold cyan][startup {self._current_step}/{self._total}][/bold cyan] {description}",
highlight=False,
)
def advance(self, description: str) -> None:
emit_boot_event(
"step_done",
phase=self._title,
step=self._current_step,
total=self._total,
message=description,
)
self._last_description = description
self._current_step = min(self._current_step + 1, self._total)
@dataclass(frozen=True)
class ModelIntegritySpec:
description: str
path: Path
required_patterns: tuple[str, ...]
alternative_required_patterns: tuple[tuple[str, ...], ...] = ()
min_total_size_bytes: int = 0
def _format_bytes(num_bytes: int) -> str:
value = float(num_bytes)
units = ["B", "KB", "MB", "GB", "TB"]
for unit in units:
if value < 1024.0 or unit == units[-1]:
return f"{value:.1f}{unit}"
value /= 1024.0
return f"{num_bytes}B"
def _find_pattern_matches(root: Path, pattern: str) -> list[Path]:
return [path for path in root.glob(pattern) if path.is_file()]
def _find_missing_patterns(root: Path, patterns: tuple[str, ...]) -> list[str]:
return [pattern for pattern in patterns if not _find_pattern_matches(root, pattern)]
def _format_alternative_patterns(pattern_groups: tuple[tuple[str, ...], ...]) -> str:
return " OR ".join(" + ".join(group) for group in pattern_groups)
def _check_model_integrity_spec(spec: ModelIntegritySpec) -> dict[str, Any]:
if not spec.path.exists() or not spec.path.is_dir():
return {
"description": spec.description,
"path": str(spec.path),
"ok": False,
"missing_patterns": [
*spec.required_patterns,
*(
[_format_alternative_patterns(spec.alternative_required_patterns)]
if spec.alternative_required_patterns
else []
),
],
"total_size_bytes": 0,
"reason": "directory_missing",
}
files = [path for path in spec.path.rglob("*") if path.is_file()]
total_size_bytes = sum(path.stat().st_size for path in files)
missing_patterns = _find_missing_patterns(spec.path, spec.required_patterns)
if not missing_patterns and spec.alternative_required_patterns:
alternative_missing_patterns = [
_find_missing_patterns(spec.path, group)
for group in spec.alternative_required_patterns
]
if all(alternative_missing_patterns):
missing_patterns = [
_format_alternative_patterns(spec.alternative_required_patterns)
]
if missing_patterns:
return {
"description": spec.description,
"path": str(spec.path),
"ok": False,
"missing_patterns": missing_patterns,
"total_size_bytes": total_size_bytes,
"reason": "required_files_missing",
}
if total_size_bytes < spec.min_total_size_bytes:
return {
"description": spec.description,
"path": str(spec.path),
"ok": False,
"missing_patterns": [],
"total_size_bytes": total_size_bytes,
"reason": "directory_too_small",
}
return {
"description": spec.description,
"path": str(spec.path),
"ok": True,
"missing_patterns": [],
"total_size_bytes": total_size_bytes,
"reason": "ok",
}
def _build_modelscope_spec(
model_id: str,
description: str,
required_patterns: tuple[str, ...],
*,
min_total_size_bytes: int,
alternative_required_patterns: tuple[tuple[str, ...], ...] = (),
) -> ModelIntegritySpec:
from ..core.config import settings
return ModelIntegritySpec(
description=description,
path=Path(settings.MODELSCOPE_PATH) / model_id,
required_patterns=required_patterns,
alternative_required_patterns=alternative_required_patterns,
min_total_size_bytes=min_total_size_bytes,
)
def _convert_ms_patterns(
patterns: tuple[str, ...],
) -> tuple[str, ...]:
return tuple(p.replace("snapshots/*/", "") for p in patterns)
def _build_qwen_spec(
model_id: str,
description: str,
required_patterns: tuple[str, ...],
*,
min_total_size_bytes: int,
alternative_required_patterns: tuple[tuple[str, ...], ...] = (),
) -> ModelIntegritySpec:
from ..core.config import settings
ms_path = Path(settings.MODELSCOPE_PATH) / model_id
ms_required = _convert_ms_patterns(required_patterns)
ms_alternative = tuple(
_convert_ms_patterns(group)
for group in alternative_required_patterns
)
return ModelIntegritySpec(
description=description,
path=ms_path,
required_patterns=ms_required,
alternative_required_patterns=ms_alternative,
min_total_size_bytes=min_total_size_bytes,
)
def _should_check_qwen_forced_aligner(
resolved_device: str,
using_cpu_qwen_rust: bool,
) -> bool:
"""Return True when startup integrity should require Qwen forced aligner files."""
from ..core.config import settings
_ = (resolved_device, using_cpu_qwen_rust)
return settings.ASR_ENABLE_WORD_TIMESTAMPS
def _build_required_model_integrity_specs() -> list[ModelIntegritySpec]:
from ..core.config import settings
from ..core.device import detect_device
from ..services.asr.manager import get_model_manager
from ..services.asr.model_capabilities import (
get_enabled_qwen_modelscope_assets,
get_runtime_required_modelscope_assets,
)
from ..services.asr.model_plan import get_runtime_model_ids
from ..services.asr.qwenasr_rust import is_qwenasr_rust_available
manager = get_model_manager()
model_ids = [item["id"] for item in manager.list_declared_entries()]
runtime_models = get_runtime_model_ids(model_ids)
resolved_device = detect_device(settings.DEVICE)
using_cpu_qwen_rust = (
resolved_device == "cpu" and is_qwenasr_rust_available()
)
specs: list[ModelIntegritySpec] = []
for asset in get_runtime_required_modelscope_assets(
include_realtime_punc=settings.ASR_ENABLE_REALTIME_PUNC,
):
specs.append(
_build_modelscope_spec(
asset.model_id,
asset.description,
asset.required_patterns,
alternative_required_patterns=asset.alternative_required_patterns,
min_total_size_bytes=asset.min_total_size_bytes,
)
)
for asset in get_enabled_qwen_modelscope_assets(
include_forced_aligner=_should_check_qwen_forced_aligner(
resolved_device=resolved_device,
using_cpu_qwen_rust=using_cpu_qwen_rust,
),
):
specs.append(
_build_qwen_spec(
asset.model_id,
asset.description,
asset.required_patterns,
alternative_required_patterns=asset.alternative_required_patterns,
min_total_size_bytes=asset.min_total_size_bytes,
)
)
return specs
def verify_required_models_integrity(use_logger: bool = True) -> dict[str, Any]:
output = logger.info if use_logger else print
specs = _build_required_model_integrity_specs()
total = len(specs)
results: list[dict[str, Any]] = []
invalid: list[dict[str, Any]] = []
if not use_logger:
output("=" * 60)
output(f"🔍 开始检查运行时模型完整性,共 {total} 个")
output("=" * 60)
for index, spec in enumerate(specs, start=1):
output(f"[{index}/{total}] 检查 {spec.description}")
result = _check_model_integrity_spec(spec)
results.append(result)
if result["ok"]:
output(
f" ✅ OK size={_format_bytes(result['total_size_bytes'])} "
f"path={result['path']}"
)
continue
invalid.append(result)
if result["reason"] == "directory_missing":
output(f" ❌ FAIL directory_missing path={result['path']}")
elif result["reason"] == "required_files_missing":
output(
f" ❌ FAIL missing={', '.join(result['missing_patterns'])} "
f"size={_format_bytes(result['total_size_bytes'])} path={result['path']}"
)
else:
output(
f" ❌ FAIL size_too_small size={_format_bytes(result['total_size_bytes'])} "
f"path={result['path']}"
)
output("=" * 60)
output(f"模型完整性检查完成: total={total} ok={total - len(invalid)} failed={len(invalid)}")
output("=" * 60)
return {
"total": total,
"results": results,
"invalid_models": invalid,
}
logger.info("开始检查运行时模型完整性: total=%s", total)
with _StartupProgress("检查运行时模型完整性", total) as progress:
for spec in specs:
progress.update(f"检查 {spec.description}")
result = _check_model_integrity_spec(spec)
results.append(result)
if not result["ok"]:
invalid.append(result)
if result["reason"] == "directory_missing":
logger.error("模型完整性检查失败: %s, reason=directory_missing, path=%s", spec.description, result["path"])
elif result["reason"] == "required_files_missing":
logger.error(
"模型完整性检查失败: %s, reason=required_files_missing, missing=%s, size=%s, path=%s",
spec.description,
", ".join(result["missing_patterns"]),
_format_bytes(result["total_size_bytes"]),
result["path"],
)
else:
logger.error(
"模型完整性检查失败: %s, reason=directory_too_small, size=%s, path=%s",
spec.description,
_format_bytes(result["total_size_bytes"]),
result["path"],
)
progress.advance(f"检查完成 {spec.description}")
logger.info(
"模型完整性检查完成: total=%s ok=%s failed=%s",
total,
total - len(invalid),
len(invalid),
)
return {
"total": total,
"results": results,
"invalid_models": invalid,
}
def preload_models() -> dict[str, Any]:
"""
预加载所有需要的模型(根据 ENABLE_* 配置过滤)
Returns:
dict: 包含加载状态的字典
"""
# 修复 CAM++ 配置文件(用于离线环境)
try:
from .download_models import fix_camplusplus_config
fix_camplusplus_config()
except Exception:
pass # 修复失败不影响启动
result: dict[str, Any] = {
"asr_models": {}, # 所有ASR模型加载状态
"vad_model": {"loaded": False, "error": None},
"speaker_diarization_model": {"loaded": False, "error": None},
}
from ..core.config import settings
from ..core.device import detect_device
# 初始化变量,避免未绑定错误
asr_device = detect_device(settings.DEVICE)
model_manager = None
# 1. 预加载所有配置的ASR模型(根据 ENABLE_* 配置过滤)
model_ids: list[str] = []
model_manager = None
try:
from ..services.asr.manager import get_model_manager
from ..services.asr.model_plan import get_runtime_model_ids
from ..services.asr.runtime import get_runtime_router
model_manager = get_model_manager()
runtime_router = get_runtime_router()
# 获取所有模型配置
all_models = model_manager.list_declared_entries()
model_ids = [m["id"] for m in all_models]
models_to_load = get_runtime_model_ids(model_ids)
if not models_to_load:
logger.warning("⚠️ 当前环境未解析出可运行的 ASR 模型")
except Exception as e:
logger.error(f"❌ 获取模型管理器失败: {e}")
models_to_load = []
runtime_router = None
total_steps = len(models_to_load) + 2
logger.info(
"开始预加载模型: declared=%s runtime=%s models=%s",
len(model_ids) if model_manager else 0,
len(models_to_load),
", ".join(models_to_load) if models_to_load else "(无)",
)
with _StartupProgress("预加载模型", total_steps) as progress:
for model_id in models_to_load:
result["asr_models"][model_id] = {"loaded": False, "error": None}
progress.update(f"加载 ASR 模型 {model_id}")
try:
if runtime_router is None:
raise RuntimeError("runtime router unavailable")
runtime_router.warmup_model(model_id)
result["asr_models"][model_id]["loaded"] = True
except Exception as e:
result["asr_models"][model_id]["error"] = str(e)
logger.error("ASR模型预加载失败: %s, error=%s", model_id, e)
progress.advance(f"已完成 ASR 模型 {model_id}")
# 2. 预加载语音活动检测模型(VAD)
progress.update("加载语音活动检测模型(VAD)")
try:
from ..services.asr.engines import get_global_vad_model
vad_model = get_global_vad_model(asr_device)
if vad_model:
result["vad_model"]["loaded"] = True
else:
result["vad_model"]["error"] = "语音活动检测模型(VAD)加载后返回None"
except Exception as e:
result["vad_model"]["error"] = str(e)
logger.error("语音活动检测模型(VAD)加载失败: %s", e)
progress.advance("已完成语音活动检测模型(VAD)")
# 5. 预加载说话人分离模型 (CAM++) - 必需模型,始终加载
progress.update("加载说话人分离模型(CAM++)")
try:
from ..utils.speaker_diarizer import get_global_diarization_pipeline
diarization_pipeline = get_global_diarization_pipeline()
if diarization_pipeline:
result["speaker_diarization_model"]["loaded"] = True
else:
result["speaker_diarization_model"]["error"] = "说话人分离模型加载后返回None"
except Exception as e:
result["speaker_diarization_model"]["error"] = str(e)
logger.error("说话人分离模型(CAM++)加载失败: %s", e)
progress.advance("已完成说话人分离模型(CAM++)")
loaded_asr_count = sum(1 for status in result["asr_models"].values() if status["loaded"])
total_asr_count = len(result["asr_models"])
extra_loaded = sum(
1
for key in ("vad_model", "speaker_diarization_model")
if result[key]["loaded"]
)
extra_failed = sum(
1
for key in ("vad_model", "speaker_diarization_model")
if result[key]["error"]
)
logger.info(
"模型预加载完成: asr=%s/%s extra_loaded=%s extra_failed=%s",
loaded_asr_count,
total_asr_count,
extra_loaded,
extra_failed,
)
return result

View File

@ -0,0 +1,697 @@
# -*- coding: utf-8 -*-
"""
说话人分离模块
基于 CAM++ 的说话人分离,用于多说话人音频分割
"""
from loguru import logger
import numpy as np
import librosa
import soundfile as sf
import tempfile
import os
import threading
from typing import Any, List, Mapping, Optional, Sequence, cast
from dataclasses import dataclass
import torch
from ..core.config import settings
from ..core.exceptions import DefaultServerErrorException
# 全局 CAM++ pipeline 缓存(单例)
_global_diarization_pipeline: Any | None = None
_diarization_pipeline_lock = threading.Lock()
_diarization_inference_semaphore = threading.BoundedSemaphore(1)
@dataclass
class SpeakerSegment:
"""说话人分段信息"""
start_ms: int
end_ms: int
speaker_id: str
audio_data: Optional[np.ndarray] = None
temp_file: Optional[str] = None
@property
def start_sec(self) -> float:
return self.start_ms / 1000.0
@property
def end_sec(self) -> float:
return self.end_ms / 1000.0
@property
def duration_ms(self) -> int:
return self.end_ms - self.start_ms
@property
def duration_sec(self) -> float:
return self.duration_ms / 1000.0
def _resolve_modelscope_device() -> str:
"""根据配置和硬件自动选择 modelscope pipeline 设备
"""
from ..core.device import detect_device
return detect_device(settings.DEVICE)
def _move_pipeline_model_to_device(pipeline_instance: Any, modelscope_device: str) -> None:
"""将 pipeline 的底层模型迁移到目标设备。"""
if hasattr(pipeline_instance, "device_name"):
pipeline_instance.device_name = modelscope_device
model = getattr(pipeline_instance, "model", None)
if model is not None and hasattr(model, "to"):
pipeline_instance.model = model.to(modelscope_device)
def _create_modelscope_pipeline(
*,
task: Any,
model: str,
modelscope_device: str,
model_revision: Optional[str] = None,
) -> Any:
"""创建 modelscope pipeline,并在需要时把底层模型迁移到目标设备。"""
from modelscope.pipelines import pipeline
pipeline_kwargs: dict[str, Any] = {
"task": task,
"model": model,
"device": modelscope_device,
}
if model_revision is not None:
pipeline_kwargs["model_revision"] = model_revision
pipeline_instance = pipeline(**pipeline_kwargs)
_move_pipeline_model_to_device(pipeline_instance, modelscope_device)
return pipeline_instance
def _enable_batched_sv(
pipeline_instance: Any,
modelscope_device: str,
max_batch_size: int = 32,
) -> Any:
"""
对说话人分离 pipeline 启用 batched SV 推理。
原始 pipeline 的 forward 方法逐个 segment 调用 sv_pipeline 提取 embedding,
这里改为将所有 segment 拼成一个 batch 一次性推理,大幅减少 GPU 调用次数。
同时将子 pipeline(sv / vad / change_locator)绑定到指定 device。
Args:
pipeline_instance: CAM++ diarization pipeline 实例
modelscope_device: 设备名称
max_batch_size: 最大批处理大小,防止 OOM
"""
if getattr(pipeline_instance, "_batched_sv_enabled", False):
return pipeline_instance
from modelscope.utils.constant import Tasks
config = getattr(pipeline_instance, "config", None)
if not isinstance(config, Mapping):
logger.warning("CAM++ pipeline 缺少可读取的 config,跳过 batched SV 优化")
return pipeline_instance
sv_model = config.get("speaker_model")
vad_model = config.get("vad_model")
change_locator = config.get("change_locator")
if isinstance(sv_model, str) and sv_model:
pipeline_instance.sv_pipeline = _create_modelscope_pipeline(
task=Tasks.speaker_verification,
model=sv_model,
modelscope_device=modelscope_device,
)
if isinstance(vad_model, str) and vad_model:
pipeline_instance.vad_pipeline = _create_modelscope_pipeline(
task=Tasks.voice_activity_detection,
model=vad_model,
modelscope_device=modelscope_device,
model_revision="v2.0.2",
)
if isinstance(change_locator, str) and change_locator:
pipeline_instance.change_locator_pipeline = _create_modelscope_pipeline(
task=Tasks.speaker_diarization,
model=change_locator,
modelscope_device=modelscope_device,
)
def batched_forward(self: Any, segments: Sequence[Sequence[Any]]) -> np.ndarray:
"""批量提取说话人 embedding,替代逐段串行推理"""
sv_model_instance = getattr(getattr(self, "sv_pipeline", None), "model", None)
emb_size = int(getattr(sv_model_instance, "emb_size", 192))
if not segments:
return np.empty((0, emb_size), dtype=np.float32)
if sv_model_instance is None:
raise RuntimeError("CAM++ sv_pipeline.model 未初始化")
all_embeddings: list[np.ndarray] = []
total_segments = len(segments)
start_idx = 0
while start_idx < total_segments:
end_idx = min(start_idx + max_batch_size, total_segments)
batch_segments = segments[start_idx:end_idx]
batch_items: list[np.ndarray] = []
for segment in batch_segments:
if len(segment) < 3:
continue
batch_items.append(np.asarray(segment[2], dtype=np.float32))
if not batch_items:
start_idx = end_idx
continue
batch = np.stack(batch_items, axis=0)
with torch.no_grad():
embeddings = sv_model_instance(
cast(Any, torch).as_tensor(batch).to(modelscope_device)
)
if isinstance(embeddings, torch.Tensor):
all_embeddings.append(embeddings.detach().cpu().numpy())
else:
all_embeddings.append(np.asarray(embeddings, dtype=np.float32))
start_idx = end_idx
if not all_embeddings:
return np.empty((0, emb_size), dtype=np.float32)
return (
np.concatenate(all_embeddings, axis=0)
if len(all_embeddings) > 1
else all_embeddings[0]
)
import types
pipeline_instance.forward = types.MethodType(batched_forward, pipeline_instance)
pipeline_instance._batched_sv_enabled = True
logger.info(
"CAM++ 说话人分离启用 batched SV: device={}, sv_device={}, vad_device={}",
modelscope_device,
getattr(getattr(pipeline_instance, "sv_pipeline", None), "device_name", "unknown"),
getattr(getattr(pipeline_instance, "vad_pipeline", None), "device_name", "unknown"),
)
return pipeline_instance
def get_global_diarization_pipeline() -> Any:
"""获取全局说话人分离 pipeline(懒加载单例)"""
global _global_diarization_pipeline
with _diarization_pipeline_lock:
if _global_diarization_pipeline is None:
try:
from modelscope.utils.constant import Tasks
from ..infrastructure.model_utils import resolve_model_path
model_id = 'iic/speech_campplus_speaker-diarization_common'
model_path = resolve_model_path(model_id)
modelscope_device = _resolve_modelscope_device()
logger.info(
"正在加载 CAM++ 说话人分离模型: {}, device={}",
model_path,
modelscope_device,
)
_global_diarization_pipeline = _create_modelscope_pipeline(
task=Tasks.speaker_diarization,
model=model_path,
modelscope_device=modelscope_device,
)
_global_diarization_pipeline = _enable_batched_sv(
_global_diarization_pipeline, modelscope_device
)
logger.info("CAM++ 模型加载成功(已启用 batched SV)")
except Exception as e:
logger.error(f"CAM++ 模型加载失败: {e}")
raise DefaultServerErrorException(f"说话人分离模型加载失败: {str(e)}")
return _global_diarization_pipeline
class SpeakerDiarizer:
"""基于 CAM++ 的说话人分离器"""
DEFAULT_MIN_SEGMENT_SEC = 1.0
DEFAULT_SAMPLE_RATE = 16000
LOW_ENERGY_SEARCH_WINDOW_MS = 10000
LOW_ENERGY_CONTEXT_MS = 160
LOW_ENERGY_STEP_MS = 20
def __init__(
self,
min_segment_sec: float = DEFAULT_MIN_SEGMENT_SEC,
):
self.min_segment_sec = min_segment_sec
self.min_segment_ms = int(min_segment_sec * 1000)
def diarize(
self, audio_path: str
) -> List[SpeakerSegment]:
"""执行说话人分离
Args:
audio_path: 音频文件路径
Returns:
原始分段列表(未合并)
"""
audio_duration_ms: Optional[int] = None
try:
try:
audio_duration_ms = int(librosa.get_duration(path=audio_path) * 1000)
except Exception:
audio_duration_ms = None
# CAM++ 对极短片段收益很低,且容易直接报 "too short"。
# 这里提前降级成单说话人,避免无意义 warning 刷屏。
if audio_duration_ms is not None and audio_duration_ms < self.min_segment_ms:
logger.debug(
"音频时长过短,跳过 CAM++ 说话人分离: duration_ms=%s < min_segment_ms=%s",
audio_duration_ms,
self.min_segment_ms,
)
return [
SpeakerSegment(
start_ms=0,
end_ms=max(audio_duration_ms, 1),
speaker_id="说话人1",
)
]
pipeline = get_global_diarization_pipeline()
logger.info(f"开始说话人分离: {audio_path}")
with _diarization_inference_semaphore:
result = pipeline(audio_path)
# 解析结果: {'text': [[start, end, speaker_id], ...]}
# pipeline 返回类型不确定,需要安全地获取 'text' 字段
if isinstance(result, dict):
raw_output = result.get('text', [])
else:
raw_output = getattr(result, 'text', []) or []
segments = []
for seg in raw_output:
if isinstance(seg, list) and len(seg) == 3:
try:
start_ms = int(float(seg[0]) * 1000)
end_ms = int(float(seg[1]) * 1000)
speaker_id = f"说话人{int(seg[2]) + 1}"
segments.append(SpeakerSegment(
start_ms=start_ms,
end_ms=end_ms,
speaker_id=speaker_id,
))
except (ValueError, TypeError) as e:
logger.warning(f"跳过格式错误的片段: {seg}, 错误: {e}")
logger.info(f"说话人分离完成,原始片段数: {len(segments)}")
# 诊断日志:打印前20个原始片段
for i, seg in enumerate(segments[:20]):
logger.debug(
f"[CAM++原始] #{i}: {seg.start_sec:.2f}-{seg.end_sec:.2f}s "
f"({seg.duration_sec:.2f}s) {seg.speaker_id}"
)
return segments
except Exception as e:
error_msg = str(e).lower()
# 音频太短时,返回默认的单说话人片段
if "too short" in error_msg:
logger.debug("CAM++ 跳过过短音频,回退单说话人片段: %s", e)
if audio_duration_ms is None:
try:
audio_duration_ms = int(librosa.get_duration(path=audio_path) * 1000)
except Exception:
audio_duration_ms = 5000
return [
SpeakerSegment(
start_ms=0,
end_ms=audio_duration_ms,
speaker_id="说话人1",
)
]
# 其他异常正常抛出
logger.error(f"说话人分离失败: {e}")
raise DefaultServerErrorException(f"说话人分离失败: {str(e)}")
def merge_consecutive_segments(
self, segments: List[SpeakerSegment]
) -> List[SpeakerSegment]:
"""合并同一说话人的连续片段"""
if not segments:
return []
# 按开始时间排序
sorted_segments = sorted(segments, key=lambda x: x.start_ms)
merged = []
current = SpeakerSegment(
start_ms=sorted_segments[0].start_ms,
end_ms=sorted_segments[0].end_ms,
speaker_id=sorted_segments[0].speaker_id,
)
for seg in sorted_segments[1:]:
if seg.speaker_id == current.speaker_id:
# 同一说话人,扩展结束时间
current.end_ms = max(current.end_ms, seg.end_ms)
else:
# 不同说话人,保存当前段,开始新段
logger.debug(
f"[合并中断] 说话人切换: {current.speaker_id} → {seg.speaker_id} "
f"在 {seg.start_sec:.2f}s,保存片段 {current.start_sec:.2f}-{current.end_sec:.2f}s"
)
merged.append(current)
current = SpeakerSegment(
start_ms=seg.start_ms,
end_ms=seg.end_ms,
speaker_id=seg.speaker_id,
)
# 保存最后一段
merged.append(current)
logger.info(f"合并同一说话人连续片段: {len(segments)} → {len(merged)}")
# 诊断日志:打印合并后的前20个片段
for i, seg in enumerate(merged[:20]):
logger.debug(
f"[合并后] #{i}: {seg.start_sec:.2f}-{seg.end_sec:.2f}s "
f"({seg.duration_sec:.2f}s) {seg.speaker_id}"
)
return merged
def merge_short_segments(
self, segments: List[SpeakerSegment]
) -> List[SpeakerSegment]:
"""智能合并短片段
策略:
1. 第一层:<10s的片段向后合并(避免孤立短片段)
2. 第二层:60s累积合并(合并连续片段)
"""
if not segments:
return []
max_segment_sec = settings.MAX_SEGMENT_SEC
# 按开始时间排序
sorted_segments = sorted(segments, key=lambda x: x.start_ms)
# 第一层:<10s累积向后合并(循环计算直到>=10s或超过60s)
merged = []
i = 0
while i < len(sorted_segments):
seg = sorted_segments[i]
# 如果>=10s,直接添加
if seg.duration_sec >= 10.0:
merged.append(seg)
i += 1
continue
# <10s,开始累积合并
current_start_ms = seg.start_ms
current_end_ms = seg.end_ms
current_duration_sec = seg.duration_sec
j = i + 1
# 累积合并,只要<10s且同说话人且不超过60s
while j < len(sorted_segments) and current_duration_sec < 10.0:
next_seg = sorted_segments[j]
if next_seg.speaker_id != seg.speaker_id:
break
new_duration = (next_seg.end_ms - current_start_ms) / 1000.0
if new_duration > max_segment_sec:
break
current_end_ms = next_seg.end_ms
current_duration_sec = new_duration
j += 1
# 创建合并后的片段
merged_seg = SpeakerSegment(
start_ms=current_start_ms,
end_ms=current_end_ms,
speaker_id=seg.speaker_id,
)
merged.append(merged_seg)
if j > i + 1:
logger.debug(
f"[第一层] {seg.speaker_id}: "
f"累积合并了 {j - i} 个片段,结果 {merged_seg.duration_sec:.1f}s"
)
i = j
# 第二层:60s累积合并
final_merged = []
i = 0
while i < len(merged):
seg = merged[i]
current_start_ms = seg.start_ms
current_end_ms = seg.end_ms
j = i + 1
# 累积合并,只要 <= 60s 且同说话人
while j < len(merged):
next_seg = merged[j]
if next_seg.speaker_id != seg.speaker_id:
break
new_duration = (next_seg.end_ms - current_start_ms) / 1000.0
if new_duration > max_segment_sec:
break
current_end_ms = next_seg.end_ms
j += 1
merged_seg = SpeakerSegment(
start_ms=current_start_ms,
end_ms=current_end_ms,
speaker_id=seg.speaker_id,
)
final_merged.append(merged_seg)
if j > i + 1:
logger.debug(
f"[第二层] {seg.speaker_id}: "
f"合并了 {j - i} 个片段"
)
i = j
return final_merged
def _find_low_energy_boundary_ms(
self,
audio_data: np.ndarray,
sample_rate: int,
lower_ms: int,
upper_ms: int,
) -> int:
lower_ms = max(0, lower_ms)
upper_ms = max(lower_ms, upper_ms)
if sample_rate <= 0 or audio_data.size == 0:
return upper_ms
context_samples = max(
1, int(sample_rate * self.LOW_ENERGY_CONTEXT_MS / 1000)
)
candidate_points = list(
range(lower_ms, upper_ms + 1, self.LOW_ENERGY_STEP_MS)
)
if not candidate_points or candidate_points[-1] != upper_ms:
candidate_points.append(upper_ms)
best_ms = upper_ms
best_energy = float("inf")
audio_length = int(audio_data.shape[0])
for candidate_ms in candidate_points:
center_sample = int(candidate_ms * sample_rate / 1000)
start_sample = max(0, center_sample - context_samples // 2)
end_sample = min(audio_length, center_sample + context_samples // 2)
if start_sample >= end_sample:
continue
window = audio_data[start_sample:end_sample]
energy = float(np.mean(np.square(window)))
if energy <= best_energy:
best_energy = energy
best_ms = candidate_ms
return best_ms
def split_long_segments(
self,
segments: List[SpeakerSegment],
audio_data: np.ndarray,
sample_rate: int,
) -> List[SpeakerSegment]:
max_segment_ms = int(settings.MAX_SEGMENT_SEC * 1000)
if max_segment_ms <= 0:
return segments
split_segments: List[SpeakerSegment] = []
for seg in segments:
if seg.duration_ms <= max_segment_ms:
split_segments.append(seg)
continue
current_start_ms = seg.start_ms
while seg.end_ms - current_start_ms > max_segment_ms:
hard_boundary_ms = current_start_ms + max_segment_ms
lower_boundary_ms = max(
current_start_ms + self.min_segment_ms,
hard_boundary_ms - self.LOW_ENERGY_SEARCH_WINDOW_MS,
)
boundary_ms = self._find_low_energy_boundary_ms(
audio_data=audio_data,
sample_rate=sample_rate,
lower_ms=lower_boundary_ms,
upper_ms=hard_boundary_ms,
)
if boundary_ms <= current_start_ms:
boundary_ms = hard_boundary_ms
split_segments.append(
SpeakerSegment(
start_ms=current_start_ms,
end_ms=boundary_ms,
speaker_id=seg.speaker_id,
)
)
current_start_ms = boundary_ms
remaining_ms = seg.end_ms - current_start_ms
if remaining_ms >= self.min_segment_ms:
split_segments.append(
SpeakerSegment(
start_ms=current_start_ms,
end_ms=seg.end_ms,
speaker_id=seg.speaker_id,
)
)
elif split_segments:
split_segments[-1].end_ms = seg.end_ms
if len(split_segments) != len(segments):
logger.info(
"Split long speaker segments by low energy: {} -> {}, max={}s",
len(segments),
len(split_segments),
settings.MAX_SEGMENT_SEC,
)
return split_segments
def split_audio_by_speakers(
self,
audio_path: str,
output_dir: Optional[str] = None,
) -> List[SpeakerSegment]:
"""完整的说话人分离流程
流程:
1. 执行CAM++说话人分离
2. 智能合并短片段(两层合并策略)
- 第一层:<10s片段累积合并
- 第二层:60s累积合并
3. 提取音频数据,保存临时文件
Args:
audio_path: 音频文件路径
output_dir: 输出目录
Returns:
SpeakerSegment 列表
"""
try:
# 1. 执行说话人分离
raw_segments = self.diarize(audio_path)
if not raw_segments:
logger.warning("说话人分离未检测到任何片段")
return []
# 2. 智能合并短片段(第一个<10s的同说话人片段向后合并)
final_segments = self.merge_short_segments(raw_segments)
# 3. Load audio before low-energy splitting and segment extraction.
logger.info("加载音频并提取片段...")
audio_data, sr = librosa.load(audio_path, sr=self.DEFAULT_SAMPLE_RATE)
sample_rate = int(sr)
final_segments = self.split_long_segments(
final_segments,
audio_data,
sample_rate,
)
logger.info(f"智能合并完成: {len(raw_segments)} → {len(final_segments)} 个片段")
output_dir = output_dir or settings.TEMP_DIR
os.makedirs(output_dir, exist_ok=True)
for idx, seg in enumerate(final_segments):
start_sample = int(seg.start_ms / 1000 * sample_rate)
end_sample = int(seg.end_ms / 1000 * sample_rate)
seg.audio_data = audio_data[start_sample:end_sample]
# 保存临时文件
temp_file = tempfile.NamedTemporaryFile(
delete=False,
suffix=".wav",
dir=output_dir,
prefix=f"{seg.speaker_id}_{idx:03d}_",
)
temp_path = temp_file.name
temp_file.close()
sf.write(temp_path, seg.audio_data, sample_rate)
seg.temp_file = temp_path
# 统计
unique_speakers = sorted(set(seg.speaker_id for seg in final_segments))
logger.info(
f"音频分割完成: {len(final_segments)} 个片段, "
f"{len(unique_speakers)} 个说话人"
)
for spk in unique_speakers:
spk_segs = [s for s in final_segments if s.speaker_id == spk]
total_time = sum(s.duration_sec for s in spk_segs)
logger.info(f" {spk}: {len(spk_segs)} 片段, {total_time:.2f}s")
return final_segments
except Exception as e:
logger.error(f"说话人分离流程失败: {e}")
raise DefaultServerErrorException(f"说话人分离失败: {str(e)}")
@staticmethod
def cleanup_segments(segments: List[SpeakerSegment]) -> None:
"""清理临时文件"""
for seg in segments:
if seg.temp_file and os.path.exists(seg.temp_file):
try:
os.remove(seg.temp_file)
except Exception as e:
logger.warning(f"清理临时文件失败: {seg.temp_file}, {e}")

View File

@ -0,0 +1,59 @@
# -*- coding: utf-8 -*-
"""
基于itntext的ITN(逆文本标准化)工具模块
使用itntext库提供高质量的中文ITN处理
"""
import logging
logger = logging.getLogger(__name__)
# itntext导入 - 延迟导入以避免初始化问题
_itntext_normalizer = None
def _get_normalizer():
"""获取itntext标准化器实例(单例模式)"""
global _itntext_normalizer
if _itntext_normalizer is None:
try:
from itntext import Normalizer
_itntext_normalizer = Normalizer(lang="zh", operator="itn")
logger.info("itntext ITN模块初始化成功")
except ImportError as e:
logger.error(f"导入itntext失败: {e}")
raise ImportError("请安装itntext库: pip install itntext")
except Exception as e:
logger.error(f"初始化itntext失败: {e}")
raise
return _itntext_normalizer
def apply_itn_to_text(text: str) -> str:
"""
对文本应用逆文本标准化(ITN)
使用itntext库进行高质量的中文ITN处理
Args:
text: 语音识别结果文本
Returns:
应用ITN后的文本
"""
if not text or not text.strip():
return text
try:
normalizer = _get_normalizer()
result = normalizer.normalize(text)
logger.debug(f"ITN处理: '{text}' -> '{result}'")
return result
except Exception as e:
logger.warning(f"ITN处理失败: {text}, 错误: {str(e)}")
return text
def normalize_asr_text(text: str, enable_itn: bool) -> str:
if not enable_itn:
return text
return apply_itn_to_text(text)

View File

@ -0,0 +1,68 @@
# crg-mcp — code-review-graph MCP 接入插件
把已部署在本机的 `D:\github-project\code-review-graph` MCP 服务接入 PI-Desktop,
让它的工具以原生 agent 工具的形式出现。**未修改该项目任何文件。**
## 接入原理
PI-Desktop 的 MCP 客户端由插件宿主承载:宿主读取插件 `manifest.json` 里的
`contributes.mcpServers`,自行拉起进程、完成 MCP 握手,并把上游每个工具发布为
`plugin_<插件id>_<server id>_<工具名>`。因此这里只需声明,无需自己写 JSON-RPC。
## 文件
| 文件 | 作用 |
| --- | --- |
| `manifest.json` | 声明 stdio MCP 服务 + `mcp.server.local` 权限 |
| `crg.cmd` | 启动包装脚本(修环境后 exec `python -m code_review_graph serve`) |
| `main.js` | 空加载器,客户端生命周期归宿主管理 |
| `dist/crg-mcp-1.0.0.piplug` | 可安装包 |
## 两个必须保留的环境修正
宿主的 `mcpProcessEnv()` 只向子进程传递
`PATH / SystemRoot / windir / TEMP / TMP / LANG`(加插件声明的 env),所以:
1. **`PYTHONPATH` 必须注入。** 项目的 editable 安装记录
`.venv\Lib\site-packages\_editable_impl_code_review_graph.pth` 指向
`D:\github_project\code-review-graph`(下划线),而项目实际位于
`D:\github-project\code-review-graph`(连字符),因此直接
`import code_review_graph` 会 `ModuleNotFoundError`。
同理 `.venv\Scripts\code-review-graph.exe` 也不能用
(`error: uv trampoline failed to canonicalize script path`)——包装脚本绕过了
这两点,改用 `python -m code_review_graph`。
2. **`USERPROFILE` / `HOMEDRIVE` / `HOMEPATH` 必须补齐。**
`code_review_graph/constants.py` 在 import 期调用 `Path.home()`,
精简环境下会抛 `RuntimeError: Could not determine home directory.`
## 暴露的工具(10 个)
默认通过 `CRG_TOOLS` 只开放审查相关的 10 个工具(上游共 30 个,全开会显著占上下文):
`build_or_update_graph_tool`、`run_postprocess_tool`、`get_minimal_context_tool`、
`get_review_context_tool`、`get_impact_radius_tool`、`query_graph_tool`、
`semantic_search_nodes_tool`、`detect_changes_tool`、`list_graph_stats_tool`、
`get_affected_flows_tool`
工具名前缀为 `plugin_crg_mcp_crg_`,例如 `plugin_crg_mcp_crg_query_graph_tool`。
### 调整暴露范围 / 目标仓库
编辑 `crg.cmd`:
- 改 `CRG_TOOLS=...`:改工具白名单(置空并删除该行 = 暴露全部 30 个)。
- 改 `CRG_REPO=...`:改被分析的仓库根目录(默认指向 code-review-graph 自身,
即已建好图的那个库)。
改完重启 PI-Desktop 生效。
## 为何用 stdio 而不是 HTTP
`serve --http`(127.0.0.1:5555/mcp)实测可用,但宿主只负责 spawn,不会托管一个
常驻服务;HTTP 需要外部进程守护,进程一掉工具就全空。stdio 由宿主拉起并在每次
调用时自动重连握手,更稳。
## 校验方式
装好后对 agent 说“用图谱统计一下仓库规模”,应命中
`plugin_crg_mcp_crg_list_graph_stats_tool` 并返回节点/边数量。

View File

@ -0,0 +1,57 @@
@echo off
chcp 65001 >nul
setlocal
rem ---------------------------------------------------------------------------
rem code-review-graph MCP launcher for PI-Desktop.
rem
rem The MCP host spawns this with a minimal environment (PATH, SystemRoot,
rem windir, TEMP, TMP, LANG plus the manifest's env block) and cwd = plugin
rem directory, with stdin/stdout used for JSON-RPC. Never write to stdout.
rem ---------------------------------------------------------------------------
if defined CRG_HOME goto have_home
set "CRG_HOME=D:\github-project\code-review-graph"
:have_home
rem code_review_graph/constants.py calls Path.home() at import time; without a
rem user profile the server dies with "Could not determine home directory.".
if defined USERPROFILE goto have_profile
set "USERPROFILE=C:\Users\%USERNAME%"
:have_profile
if defined HOMEDRIVE goto have_hd
set "HOMEDRIVE=C:"
:have_hd
if defined HOMEPATH goto have_hp
set "HOMEPATH=\Users\%USERNAME%"
:have_hp
if defined APPDATA set "APPDATA=%USERPROFILE%\AppData\Roaming"
if defined LOCALAPPDATA set "LOCALAPPDATA=%USERPROFILE%\AppData\Local"
set "PYTHONUTF8=1"
set "PYTHONIOENCODING=utf-8"
rem The editable install's .pth points at D:\github_project\... (underscore)
rem while the checkout lives at D:\github-project\... (hyphen), so the package
rem is only importable with the checkout explicitly on sys.path.
set "PYTHONPATH=%CRG_HOME%"
rem The venv's code-review-graph.exe is a broken uv trampoline
rem ("failed to canonicalize script path"), so prefer a working interpreter.
set "CRG_PY=%CRG_HOME%\.venv\Scripts\python.exe"
if exist "%CRG_PY%" goto py_ready
set "CRG_PY=python"
:py_ready
rem Upstream ships 30 tools; keep the surface small unless overridden.
if defined CRG_TOOLS goto tools_ready
set "CRG_TOOLS=build_or_update_graph_tool,run_postprocess_tool,get_minimal_context_tool,get_review_context_tool,get_impact_radius_tool,query_graph_tool,semantic_search_nodes_tool,detect_changes_tool,list_graph_stats_tool,get_affected_flows_tool"
:tools_ready
if defined CRG_REPO goto repo_ready
set "CRG_REPO=%CRG_HOME%"
:repo_ready
"%CRG_PY%" -m code_review_graph serve --repo "%CRG_REPO%"
endlocal

Binary file not shown.

View File

@ -0,0 +1,18 @@
/**
* crg-mcp — thin loader for the code-review-graph MCP bridge.
*
* All of the wiring lives in manifest.json under `contributes.mcpServers`:
* PI-Desktop spawns `crg.cmd` (stdio MCP), performs the handshake, and
* publishes every upstream tool as `plugin_crg_mcp_crg_<tool>`. Nothing has
* to be registered from here — the host owns the client, the retries, and the
* tool lifecycle. This module only keeps the plugin loadable and offers a
* place for future local helpers.
*/
async function onLoad() {
// The MCP client is owned by the host; no tool registration needed.
}
async function onUnload() {}
module.exports = { onLoad, onUnload };

View File

@ -0,0 +1,33 @@
{
"schemaVersion": 1,
"id": "crg-mcp",
"name": "Code Review Graph MCP",
"version": "1.0.0",
"description": "Bridges the locally deployed code-review-graph MCP server (D:\\github-project\\code-review-graph) into the agent as native tools.",
"main": "main.js",
"contributes": {
"mcpServers": [
{
"id": "crg",
"label": "Code Review Graph",
"transport": "stdio",
"command": "crg.cmd",
"env": {
"PYTHONPATH": "D:\\github-project\\code-review-graph",
"USERPROFILE": "C:\\Users\\admin",
"HOMEDRIVE": "C:",
"HOMEPATH": "\\Users\\admin"
}
}
]
},
"permissions": [
"mcp.server.local"
],
"engines": {
"piDesktop": ">=0.1.0"
},
"activationEvents": [
"onStartup"
]
}

View File

@ -0,0 +1,96 @@
<#
Self-check for the crg-mcp plugin.
Drives crg.cmd exactly the way the PI-Desktop MCP host does — `cmd /c crg.cmd`
with piped stdio and cwd = this folder — then reports the MCP handshake, the
tool list, and one real tool call. Run it from PowerShell:
powershell -NoProfile -ExecutionPolicy Bypass -File .\selfcheck.ps1
#>
$ErrorActionPreference = 'Stop'
$dir = Split-Path -Parent $MyInvocation.MyCommand.Path
# Mirror the host's minimal environment plus the manifest's env block.
foreach ($k in 'PATH', 'SystemRoot', 'TEMP', 'TMP') {
if (-not (Test-Path "Env:$k")) { Write-Warning "missing $k in ambient env" }
}
$psi = New-Object System.Diagnostics.ProcessStartInfo
$psi.FileName = 'cmd.exe'
$psi.Arguments = '/c crg.cmd'
$psi.WorkingDirectory = $dir
$psi.RedirectStandardInput = $true
$psi.RedirectStandardOutput = $true
$psi.RedirectStandardError = $true
$psi.UseShellExecute = $false
$psi.StandardOutputEncoding = [System.Text.Encoding]::UTF8
$proc = [System.Diagnostics.Process]::Start($psi)
function Send($obj) {
$proc.StandardInput.WriteLine(($obj | ConvertTo-Json -Compress -Depth 8))
$proc.StandardInput.Flush()
}
function ReadLine([int]$waitSeconds = 25) {
$task = $proc.StandardOutput.ReadLineAsync()
if ($task.Wait([TimeSpan]::FromSeconds($waitSeconds))) { return $task.Result }
return $null
}
Send @{
jsonrpc = '2.0'; id = 1; method = 'initialize'
params = @{
protocolVersion = '2025-06-18'
capabilities = @{}
clientInfo = @{ name = 'crg-selfcheck'; version = '1' }
}
}
$initLine = ReadLine 40
if (-not $initLine) {
Write-Host 'HANDSHAKE FAILED: no stdout from crg.cmd' -ForegroundColor Red
Write-Host '--- stderr ---'
Write-Host $proc.StandardError.ReadToEnd()
try { $proc.Kill() } catch { }
exit 1
}
$init = $initLine | ConvertFrom-Json
Write-Host ("HANDSHAKE OK server={0} {1}" -f $init.result.serverInfo.name, $init.result.serverInfo.version) -ForegroundColor Green
Send @{ jsonrpc = '2.0'; method = 'notifications/initialized'; params = @{} }
Send @{ jsonrpc = '2.0'; id = 2; method = 'tools/list'; params = @{} }
$toolsLine = ReadLine
if (-not $toolsLine) {
Write-Host 'tools/list returned nothing' -ForegroundColor Red
try { $proc.Kill() } catch { }
exit 1
}
$tools = ($toolsLine | ConvertFrom-Json).result.tools
Write-Host ("TOOLS: {0}" -f $tools.Count) -ForegroundColor Green
foreach ($t in $tools) { Write-Host (" plugin_crg_mcp_crg_{0}" -f $t.name) }
Send @{
jsonrpc = '2.0'; id = 3; method = 'tools/call'
params = @{ name = 'list_graph_stats_tool'; arguments = @{} }
}
$callLine = ReadLine 40
if ($callLine) {
$call = $callLine | ConvertFrom-Json
if ($call.result) {
$text = $call.result.content[0].text
Write-Host 'TOOL CALL OK' -ForegroundColor Green
Write-Host (' ' + ($text -split "`n")[0..3] -join ' | ')
}
else {
Write-Host ("TOOL CALL ERROR: {0}" -f ($call | ConvertTo-Json -Compress -Depth 6)) -ForegroundColor Red
}
}
else {
Write-Host 'TOOL CALL: no response' -ForegroundColor Red
}
try { $proc.Kill() } catch { }

View File

@ -0,0 +1,35 @@
name: qwen3-asr
services:
qwen3-asr:
image: ${ASR_IMAGE:-unis/qwen3-asr:cpu-latest}
container_name: qwen3-asr-cpu
ports:
- "${NGINX_PORT:-17003}:8000"
volumes:
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
environment:
ACCELERATOR: cpu
API_KEY: ${API_KEY:-}
MODELS_DIR: /app/models
DATA_DIR: /app/data
TEMP_DIR: /app/data/temp
LOG_FILE: /app/data/logs/qwen3-asr.log
TASK_STATE_DIR: /app/data/tasks
TASK_RETENTION_HOURS: ${TASK_RETENTION_HOURS:-24}
MODELSCOPE_CACHE: /app
MODELSCOPE_PATH: /app/models
QWEN3_ASR_MODEL: ${QWEN3_ASR_MODEL:-}
SPEAKER_DB_ENABLED: ${SPEAKER_DB_ENABLED:-true}
DB_HOST: ${DB_HOST:-127.0.0.1}
DB_PORT: ${DB_PORT:-5432}
DB_USER: ${DB_USER:-postgres}
DB_PASSWORD: ${DB_PASSWORD:-postgres}
DB_NAME: ${DB_NAME:-asr_db}
SV_MODEL: ${SV_MODEL:-iic/speech_campplus_sv_zh-cn_16k-common}
SV_THRESHOLD: ${SV_THRESHOLD:-0.6}
REALTIME_SESSION_RESUME_TTL_SEC: ${REALTIME_SESSION_RESUME_TTL_SEC:-120}
NGINX_RATE_LIMIT_RPS: ${NGINX_RATE_LIMIT_RPS:-0}
NGINX_RATE_LIMIT_BURST: ${NGINX_RATE_LIMIT_BURST:-0}
restart: unless-stopped

View File

@ -0,0 +1,50 @@
name: qwen3-asr-iluvatar
services:
qwen3-asr:
image: ${ASR_IMAGE:-unis/qwen3-asr:iluvatar-latest}
container_name: qwen3-asr-iluvatar
network_mode: host
pid: host
ipc: host
privileged: true
cap_add:
- ALL
volumes:
- ${ILUVATAR_USR_SRC:-/usr/src}:/usr/src
- ${ILUVATAR_LIB_MODULES:-/lib/modules}:/lib/modules
- ${ILUVATAR_DEV:-/dev}:/dev
- ${ILUVATAR_HOME:-/home}:/home
- ${ILUVATAR_DATA:-/data}:/data
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
environment:
ACCELERATOR: iluvatar
DEVICE: ${DEVICE:-auto}
PORT: ${NGINX_PORT:-17003}
API_KEY: ${API_KEY:-}
MODELS_DIR: /app/models
DATA_DIR: /app/data
TEMP_DIR: /app/data/temp
LOG_FILE: /app/data/logs/qwen3-asr.log
TASK_STATE_DIR: /app/data/tasks
TASK_RETENTION_HOURS: ${TASK_RETENTION_HOURS:-24}
MODELSCOPE_CACHE: /app
MODELSCOPE_PATH: /app/models
QWEN3_ASR_MODEL: ${QWEN3_ASR_MODEL:-}
ASR_DEPLOY_TOPOLOGY: ${ASR_DEPLOY_TOPOLOGY:-isolated}
QWEN_GPU_MEMORY_UTILIZATION: ${QWEN_GPU_MEMORY_UTILIZATION:-}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
SPEAKER_DB_ENABLED: ${SPEAKER_DB_ENABLED:-true}
DB_HOST: ${DB_HOST:-127.0.0.1}
DB_PORT: ${DB_PORT:-5432}
DB_USER: ${DB_USER:-postgres}
DB_PASSWORD: ${DB_PASSWORD:-postgres}
DB_NAME: ${DB_NAME:-asr_db}
SV_MODEL: ${SV_MODEL:-iic/speech_campplus_sv_zh-cn_16k-common}
SV_THRESHOLD: ${SV_THRESHOLD:-0.6}
REALTIME_SESSION_RESUME_TTL_SEC: ${REALTIME_SESSION_RESUME_TTL_SEC:-120}
ASR_VISIBLE_DEVICES: ${ASR_VISIBLE_DEVICES:-0}
NGINX_RATE_LIMIT_RPS: ${NGINX_RATE_LIMIT_RPS:-0}
NGINX_RATE_LIMIT_BURST: ${NGINX_RATE_LIMIT_BURST:-0}
restart: unless-stopped

View File

@ -0,0 +1,47 @@
name: qwen3-asr-metax
services:
qwen3-asr:
image: ${ASR_IMAGE:-unis/qwen3-asr:metax-latest}
container_name: qwen3-asr-metax
network_mode: host
pid: host
ipc: host
privileged: true
cap_add:
- ALL
volumes:
- ${METAX_DEV:-/dev}:/dev
- ${METAX_DRIVER_DIR:-/opt/mxdriver}:/opt/mxdriver:ro
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
environment:
ACCELERATOR: metax
DEVICE: ${DEVICE:-auto}
PORT: ${NGINX_PORT:-17003}
API_KEY: ${API_KEY:-}
MODELS_DIR: /app/models
DATA_DIR: /app/data
TEMP_DIR: /app/data/temp
LOG_FILE: /app/data/logs/qwen3-asr.log
TASK_STATE_DIR: /app/data/tasks
TASK_RETENTION_HOURS: ${TASK_RETENTION_HOURS:-24}
MODELSCOPE_CACHE: /app
MODELSCOPE_PATH: /app/models
QWEN3_ASR_MODEL: ${QWEN3_ASR_MODEL:-}
ASR_DEPLOY_TOPOLOGY: ${ASR_DEPLOY_TOPOLOGY:-isolated}
QWEN_GPU_MEMORY_UTILIZATION: ${QWEN_GPU_MEMORY_UTILIZATION:-}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
SPEAKER_DB_ENABLED: ${SPEAKER_DB_ENABLED:-true}
DB_HOST: ${DB_HOST:-127.0.0.1}
DB_PORT: ${DB_PORT:-5432}
DB_USER: ${DB_USER:-postgres}
DB_PASSWORD: ${DB_PASSWORD:-postgres}
DB_NAME: ${DB_NAME:-asr_db}
SV_MODEL: ${SV_MODEL:-iic/speech_campplus_sv_zh-cn_16k-common}
SV_THRESHOLD: ${SV_THRESHOLD:-0.6}
REALTIME_SESSION_RESUME_TTL_SEC: ${REALTIME_SESSION_RESUME_TTL_SEC:-120}
ASR_VISIBLE_DEVICES: ${ASR_VISIBLE_DEVICES:-0}
NGINX_RATE_LIMIT_RPS: ${NGINX_RATE_LIMIT_RPS:-0}
NGINX_RATE_LIMIT_BURST: ${NGINX_RATE_LIMIT_BURST:-0}
restart: unless-stopped

View File

@ -0,0 +1,50 @@
name: qwen3-asr-mthreads
services:
qwen3-asr:
image: ${ASR_IMAGE:-unis/qwen3-asr:mthreads-latest}
container_name: qwen3-asr-mthreads
network_mode: host
pid: host
ipc: host
privileged: true
cap_add:
- ALL
volumes:
- ${MTHREADS_DEV:-/dev}:/dev
- ${MTHREADS_USR_SRC:-/usr/src}:/usr/src
- ${MTHREADS_LIB_MODULES:-/lib/modules}:/lib/modules
- ${MTHREADS_HOME:-/home}:/home
- ${MTHREADS_DATA:-/data}:/data
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
environment:
ACCELERATOR: mthreads
DEVICE: ${DEVICE:-auto}
PORT: ${NGINX_PORT:-17003}
API_KEY: ${API_KEY:-}
MODELS_DIR: /app/models
DATA_DIR: /app/data
TEMP_DIR: /app/data/temp
LOG_FILE: /app/data/logs/qwen3-asr.log
TASK_STATE_DIR: /app/data/tasks
TASK_RETENTION_HOURS: ${TASK_RETENTION_HOURS:-24}
MODELSCOPE_CACHE: /app
MODELSCOPE_PATH: /app/models
QWEN3_ASR_MODEL: ${QWEN3_ASR_MODEL:-}
ASR_DEPLOY_TOPOLOGY: ${ASR_DEPLOY_TOPOLOGY:-isolated}
QWEN_GPU_MEMORY_UTILIZATION: ${QWEN_GPU_MEMORY_UTILIZATION:-}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
SPEAKER_DB_ENABLED: ${SPEAKER_DB_ENABLED:-true}
DB_HOST: ${DB_HOST:-127.0.0.1}
DB_PORT: ${DB_PORT:-5432}
DB_USER: ${DB_USER:-postgres}
DB_PASSWORD: ${DB_PASSWORD:-postgres}
DB_NAME: ${DB_NAME:-asr_db}
SV_MODEL: ${SV_MODEL:-iic/speech_campplus_sv_zh-cn_16k-common}
SV_THRESHOLD: ${SV_THRESHOLD:-0.6}
REALTIME_SESSION_RESUME_TTL_SEC: ${REALTIME_SESSION_RESUME_TTL_SEC:-120}
ASR_VISIBLE_DEVICES: ${ASR_VISIBLE_DEVICES:-0}
NGINX_RATE_LIMIT_RPS: ${NGINX_RATE_LIMIT_RPS:-0}
NGINX_RATE_LIMIT_BURST: ${NGINX_RATE_LIMIT_BURST:-0}
restart: unless-stopped

40
docker-compose.yml 100644
View File

@ -0,0 +1,40 @@
name: qwen3-asr
services:
qwen3-asr:
image: ${ASR_IMAGE:-unis/qwen3-asr:gpu-latest}
container_name: qwen3-asr
ports:
- "${NGINX_PORT:-17003}:8000"
volumes:
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
runtime: nvidia
environment:
ACCELERATOR: nvidia
API_KEY: ${API_KEY:-}
MODELS_DIR: /app/models
DATA_DIR: /app/data
TEMP_DIR: /app/data/temp
LOG_FILE: /app/data/logs/qwen3-asr.log
TASK_STATE_DIR: /app/data/tasks
TASK_RETENTION_HOURS: ${TASK_RETENTION_HOURS:-24}
MODELSCOPE_CACHE: /app
MODELSCOPE_PATH: /app/models
QWEN3_ASR_MODEL: ${QWEN3_ASR_MODEL:-}
ASR_DEPLOY_TOPOLOGY: ${ASR_DEPLOY_TOPOLOGY:-isolated}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
SPEAKER_DB_ENABLED: ${SPEAKER_DB_ENABLED:-true}
DB_HOST: ${DB_HOST:-127.0.0.1}
DB_PORT: ${DB_PORT:-5432}
DB_USER: ${DB_USER:-postgres}
DB_PASSWORD: ${DB_PASSWORD:-postgres}
DB_NAME: ${DB_NAME:-asr_db}
SV_MODEL: ${SV_MODEL:-iic/speech_campplus_sv_zh-cn_16k-common}
SV_THRESHOLD: ${SV_THRESHOLD:-0.6}
REALTIME_SESSION_RESUME_TTL_SEC: ${REALTIME_SESSION_RESUME_TTL_SEC:-120}
NVIDIA_VISIBLE_DEVICES: all
ASR_VISIBLE_DEVICES: ${ASR_VISIBLE_DEVICES:-0}
NGINX_RATE_LIMIT_RPS: ${NGINX_RATE_LIMIT_RPS:-0}
NGINX_RATE_LIMIT_BURST: ${NGINX_RATE_LIMIT_BURST:-0}
restart: unless-stopped

606
docs/README_zh.md 100644
View File

@ -0,0 +1,606 @@
<div align="center">
<h1>Qwen3-ASR</h1>
<h3>开箱即用的本地私有化部署语音识别服务</h3>
以 [Qwen3-ASR](https://github.com/QwenLM/Qwen3-ASR) 为核心的语音识别 API 服务,提供 NVIDIA CUDA vLLM、沐曦 MACA vLLM 与 CPU Rust 后端,兼容阿里云语音 API 和 OpenAI Audio API,并保留 Paraformer realtime WebSocket 能力。
---
![Static Badge](https://img.shields.io/badge/Python-3.10+-blue?logo=python)
![Static Badge](https://img.shields.io/badge/Torch-2.11.0-%23EE4C2C?logo=pytorch&logoColor=white)
![Static Badge](https://img.shields.io/badge/CUDA-13.0_default-%2376B900?logo=nvidia&logoColor=white)
</div>
## 在线演示站点
- **在线体验**: https://asr.vect.one
## 演示
[![演示](../demo/demo.png)](https://media.cdn.vect.one/qwenasr_client_demo.mp4)
## Release 1.0.1
> `v1.0.1` 是当前补丁版本。`v1.0.0` 相对于早期 `main` 分支引入了一轮大规模 breaking refactor。
> 如果你是从 `main` 升级过来,请先阅读 release 说明,再决定是否沿用旧的部署与运行时假设。
>
> 关键 breaking changes:
> - Python 依赖管理已经切到 `uv`(`pyproject.toml` + `uv.lock`),`requirements*.txt` 已移除
> - 运行时栈改成 `NVIDIA/沐曦 GPU -> vLLM`、`CPU/macOS -> vendored QwenASR Rust`
> - `MLX` / Apple Silicon GPU 路径已移除,`mps` 会归一化到 `cpu`
> - macOS / Apple Silicon 现在默认总是 `qwen3-asr-0.6b`,可通过 `QWEN3_ASR_MODEL` 覆盖
> - `ENABLED_MODELS` 已移除
## 主要特性
- **混合运行时栈** - 离线推理由自动选择的 Qwen3-ASR 提供,WebSocket 流式由 Paraformer realtime 能力提供
- **说话人分离** - 基于 CAM++ 模型自动识别多说话人,返回说话人标记
- **OpenAI API 兼容** - 支持 `/v1/audio/transcriptions` 端点,可直接使用 OpenAI SDK
- **阿里云 API 兼容** - 支持阿里云语音识别 RESTful API 和 WebSocket 流式协议
- **WebSocket 流式识别** - 支持实时流式语音识别,低延迟
- **智能远场过滤** - 流式 ASR 自动过滤远场声音和环境音,减少误触发
- **智能音频分段** - 基于 VAD 的贪婪合并算法,自动切分长音频,避免包含过长静音
- **GPU 批处理加速** - 支持批量推理,比逐个处理快 2-3 倍
- **资源感知运行时** - 根据当前机器资源自动选择合适的 Qwen3-ASR 模型
## 致谢
- [Qwen3-ASR](https://github.com/QwenLM/Qwen3-ASR) 提供官方模型与多模态 / vLLM 使用方式
- [QwenASR](https://github.com/huanglizhuo/QwenASR) 提供本项目 vendored 的 CPU Rust backend
## 快速部署
### 1. Docker 部署(推荐)
```bash
# 复制并编辑配置
cp .env.example .env
# 编辑 .env 设置 API_KEY(可选)
# Compose 默认挂载:
# /opt/dep/asr/models -> /app/models
# /opt/dep/asr/data -> /app/data
# /opt/dep/asr/data/logs、temp、tasks 都在 data 挂载内
# 启动服务(NVIDIA GPU 版本)
docker-compose up -d
# 或沐曦 GPU 版本
docker-compose -f docker-compose-metax.yml up -d
# 或天数 GPU 版本
docker-compose -f docker-compose-iluvatar.yml up -d
# 或摩尔线程 / MUSA GPU 版本
docker-compose -f docker-compose-mthreads.yml up -d
# 或 CPU 版本
docker-compose -f docker-compose-cpu.yml up -d
# NVIDIA 多卡自动模式(每张可见卡自动拉起 1 个实例)
CUDA_VISIBLE_DEVICES=0,1,2,3 docker-compose up -d
# 沐曦多卡自动模式
METAX_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-metax.yml up -d
# 天数多卡自动模式
ILUVATAR_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-iluvatar.yml up -d
# 摩尔线程 / MUSA 多卡自动模式
MTHREADS_VISIBLE_DEVICES=0,1 docker-compose -f docker-compose-mthreads.yml up -d
```
服务访问地址:
- **API 端点**: `http://localhost:17003`
- **API 文档**: `http://localhost:17003/docs`
可选的内置限流参数:
- `NGINX_RATE_LIMIT_RPS`(全局每秒请求上限,`0` 表示关闭)
- `NGINX_RATE_LIMIT_BURST`(全局突发请求数,`0` 时自动使用 RPS)
**docker run 方式(替代):**
```bash
# NVIDIA GPU 版本
docker run -d --name qwen3-asr \
--gpus all \
-p 17003:8000 \
-e ACCELERATOR=nvidia \
-e CUDA_VISIBLE_DEVICES=0,1,2,3 \
-e API_KEY=your_api_key \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:gpu-latest
# 沐曦 GPU 版本
docker run -d --name qwen3-asr-metax \
--privileged \
--network=host \
--pid=host \
--ipc=host \
-v /dev:/dev \
-v /opt/mxdriver:/opt/mxdriver:ro \
-e ACCELERATOR=metax \
-e PORT=17003 \
-e METAX_VISIBLE_DEVICES=0 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:metax-latest
# CPU 版本
docker run -d --name qwen3-asr \
-p 17003:8000 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
unis/qwen3-asr:cpu-latest
```
默认推荐将宿主机目录统一挂载到 `/opt/dep/asr` 下,模型目录结构如下:
```text
/opt/dep/asr/models/
Qwen/
iic/
damo/
```
如果你希望改成自定义目录,也可以在启动前设置:
```bash
export MODEL_STORAGE_DIR=/data/qwen3-asr-models
export DATA_STORAGE_DIR=/data/qwen3-asr-data
```
> **注意**: NVIDIA GPU 镜像默认使用 CUDA 13.0/cu130,并固定 `torch 2.11.0` + `vllm 0.20.0`。
> 开发者可通过 Docker build args 自行构建 CUDA 12.6、CUDA 13.0 或其他后端组合。
> 沐曦镜像使用 `Dockerfile.metax` 基于沐曦官方 vLLM 镜像融合本项目。现场部署使用 host network、privileged,并挂载 `/dev` 与 `/opt/mxdriver`,确保 `mx-smi` 查询和沐曦 PyTorch 运行时都能初始化设备。
> 当前 CPU 镜像已通过内置 QwenASR Rust backend 支持 `qwen3-asr-0.6b`。默认 CPU 镜像使用可分发 Rust 构建目标;只有自建且构建机/部署机 CPU 同构时才建议设置 `QWENASR_RUST_TARGET_CPU=native`。
> CUDA vLLM 与 CPU Rust 路径下,`word_timestamps=true` 都会自动调用 forced aligner;当前实际后端为 `CUDA -> vLLM`、`CPU/macOS -> vendored QwenASR Rust`。
> Apple Silicon 上的 Qwen3-ASR 现已统一走 Rust CPU backend。
> `start.py` 现在会强制把 vLLM 多进程方式设为 `spawn`,避免 CUDA 在 fork 子进程中重复初始化导致启动失败。
**自定义 GPU 后端构建:**
```bash
# 默认 GPU 构建:CUDA 13.0 / PyTorch cu130
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu .
# CUDA 12.6 构建,用于旧部署环境
docker build -t qwen3-asr:gpu-cu126 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda12.6-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu126 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-12-6 \
--build-arg TORCH_CUDA_ARCH_LIST="8.0;8.6;8.9" \
.
# CUDA 13.0 构建,用于需要 CUDA 13 工具链的环境
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda13.0-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu130 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-13-0 \
--build-arg TORCH_CUDA_ARCH_LIST="12.0+PTX" \
.
# 沐曦构建:基于沐曦官方 vLLM 镜像融合本项目
./scripts/package_vendor_gpu_image.sh \
--vendor metax \
--base-image <沐曦官方vLLM镜像名> \
-v n260-3.7.0.38
# 天数构建:基于天数官方 vLLM 镜像融合本项目
docker pull registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
./scripts/package_vendor_gpu_image.sh \
--vendor iluvatar \
--base-image registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5
# 摩尔线程构建:基于摩尔线程官方 MUSA vLLM 镜像融合本项目
docker pull registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
./scripts/package_vendor_gpu_image.sh \
--vendor mthreads \
--base-image registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519 \
-v s4000_4.3.5_d0519
```
沐曦 GPU 国产化离线交付请优先参考 [沐曦 GPU 国产化离线部署指南](./metax_offline_deployment.md)。
天数 GPU 国产化离线交付请优先参考 [天数 GPU 国产化离线部署指南](./iluvatar_offline_deployment.md)。
摩尔线程 GPU 国产化离线交付请优先参考 [摩尔线程 GPU 国产化离线部署指南](./mthreads_offline_deployment.md)。
**内网部署**:现在可以直接生成一个带时间戳和 CPU/GPU 标识的离线交付目录,里面包含镜像包、compose、`.env` 模板、目录初始化脚本和使用说明。离线导出脚本使用普通 `docker build` + `docker save`,不依赖 `buildx`:
```bash
# 1. 生成离线交付目录
./export_offline_bundle.sh --type gpu
# 或
./export_offline_bundle.sh --type cpu
# 或沐曦 GPU
./export_offline_bundle.sh \
--type metax \
--metax-base cr.metax-tech.com/public-ai-release/maca/vllm-metax:0.17.0-maca.ai3.5.3.307-torch2.8-py312-ubuntu22.04-amd64 \
--skip-models
# 或天数 GPU
./export_offline_bundle.sh --type iluvatar --iluvatar-base registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
# 或摩尔线程 GPU
./export_offline_bundle.sh --type mthreads --mthreads-base registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
# 或一次同时打包 GPU + CPU
./export_offline_bundle.sh --type all
# 2. 单独准备模型,不删除已有模型文件
./scripts/download-models.sh --models-dir /opt/dep/asr/models
# 3. 把交付目录复制到内网服务器
scp -r build-file/<时间戳>-all user@server:/opt/dep/asr/
# 4. 在内网服务器上导入并启动
cd /opt/dep/asr/<时间戳>-all
./init_host_dirs.sh
gunzip -c qwen3-asr-gpu-<时间戳>-amd64.tar.gz | docker load
gunzip -c qwen3-asr-cpu-<时间戳>-amd64.tar.gz | docker load
# NVIDIA GPU
docker compose up -d
# 或沐曦 GPU
# docker compose -f docker-compose-metax.yml up -d
# 或天数 GPU
# docker compose -f docker-compose-iluvatar.yml up -d
# 或摩尔线程 GPU
# docker compose -f docker-compose-mthreads.yml up -d
# 或 CPU
# docker compose -f docker-compose-cpu.yml up -d
```
> 详细部署说明请查看 [部署指南](./deployment.md)
### 本地开发
**系统要求:**
- Python 3.10+
- 默认 GPU 镜像要求 CUDA 13.0+;CUDA 12.6 / 13.0 可通过 Docker build args 自行构建
- FFmpeg (音频格式转换)
**安装步骤:**
运行时依赖现在改成“根目录默认 GPU,CPU 单独特化环境”:
| 模式 | 命令 | 说明 |
|------|------|------|
| NVIDIA GPU(默认) | `uv sync` 或 `./scripts/sync_gpu_env.sh` | 同步根目录 [pyproject.toml](/opt/qwen3-asr/pyproject.toml) 和 [uv.lock](/opt/qwen3-asr/uv.lock) 到 `.venv`,包含 CUDA 13.0/cu130 `torch 2.11.0` / `torchaudio 2.11.0` / `torchvision 0.26.0` / `vllm 0.20.0` |
| 沐曦 GPU | `./scripts/sync_metax_env.sh` | 同步 [environments/metax/pyproject.toml](/opt/qwen3-asr/environments/metax/pyproject.toml) 的公共依赖;可选 GPU 栈安装默认从沐曦 MACA PyPI 源按 `--no-deps` 安装 |
| 天数 GPU | `./scripts/sync_iluvatar_env.sh` | 同步 [environments/iluvatar/pyproject.toml](/opt/qwen3-asr/environments/iluvatar/pyproject.toml) 的公共依赖;GPU 栈建议来自天数官方 vLLM 镜像 |
| 摩尔线程 GPU | `./scripts/sync_mthreads_env.sh` | 同步 [environments/mthreads/pyproject.toml](/opt/qwen3-asr/environments/mthreads/pyproject.toml) 的公共依赖;GPU 栈建议来自摩尔线程官方 MUSA vLLM 镜像 |
| CPU(特化) | `./scripts/sync_cpu_env.sh` | 同步 [environments/cpu/pyproject.toml](/opt/qwen3-asr/environments/cpu/pyproject.toml) 对应的 CPU lock 到 `.venv` |
| 自动 | `./scripts/sync_accel_env.sh` | 有 `mx-smi` 时选择沐曦,有 `ixsmi` 时选择天数,有 `mthreads-gmi` 时选择摩尔线程,有 `nvidia-smi` 时选择 NVIDIA,否则选择 CPU |
```bash
# 克隆项目
cd qwen3-asr
# 安装依赖(Linux/NVIDIA CUDA)
uv sync
# 启动服务
source .venv/bin/activate
python start.py
```
沐曦本地开发:
```bash
./scripts/sync_metax_env.sh
source .venv/bin/activate
ACCELERATOR=metax python start.py
```
macOS / Apple Silicon 本地开发:
```bash
./scripts/sync_cpu_env.sh
source .venv/bin/activate
python start.py
```
## 当前运行时默认值
当前主线代码的运行时行为如下:
- `ACCELERATOR=auto` 会优先识别 `mx-smi` 上报的沐曦设备,其次识别 `ixsmi` 上报的天数设备,再识别 `mthreads-gmi` 上报的摩尔线程设备,再识别 NVIDIA CUDA,否则回落 CPU
- `DEVICE=auto`
- NVIDIA/沐曦/天数 GPU 时解析为 `cuda:0`
- 否则解析为 `cpu`
- `DEVICE=mps` 会直接归一化为 `cpu`
- `Linux + NVIDIA CUDA` 使用官方 `vLLM`
- `Linux + 沐曦 MACA` 使用沐曦兼容 PyTorch/vLLM 运行栈
- `Linux + 天数` 使用天数官方 vLLM 镜像运行栈
- `Linux + CPU` 使用 vendored `QwenASR` Rust
- `macOS / Apple Silicon` 也使用 vendored `QwenASR` Rust
- macOS / Apple Silicon 默认总是 `qwen3-asr-0.6b`
- 在 macOS 上,只有设置 `QWEN3_ASR_MODEL=qwen3-asr-1.7b` 时才会使用 `qwen3-asr-1.7b`
- `word_timestamps=true` 在当前离线 CUDA 与 CPU Rust 路径下可用
- WebSocket 流式路径当前不返回词级时间戳
- CAM++ 说话人分离仍然必须保留,并继续跟随 `DEVICE`;在 CPU 上的主要热点仍是 speaker verification embedding
## API 接口
### OpenAI 兼容接口
| 端点 | 方法 | 功能 |
| ---------------------------- | ---- | ----------------------- |
| `/v1/audio/transcriptions` | POST | 音频转写(OpenAI 兼容) |
| `/v1/models` | GET | 离线模型列表 |
**请求参数:**
| 参数 | 类型 | 默认值 | 说明 |
| ------------------------------ | ------ | --------------------- | ------------------------------------- |
| `file` | file | 提供时优先使用 | 音频/视频文件 |
| `audio_address` | string | 可选 | 音频/视频文件 URL(HTTP/HTTPS)、`file://` 或服务端本地路径;若同时提供 `file`,则忽略 |
| `language` | string | 自动检测 | 语言代码 (zh/en/ja) |
| `enable_speaker_diarization` | bool | `true` | 启用说话人分离 |
| `enable_speaker_identification` | bool | `true` | 说话人分离开启时匹配已注册声纹库 |
| `enable_text_cleanup` | bool | `true` | 启用文本去重、跨段重叠裁剪和口头语清理 |
| `word_timestamps` | bool | `false` | 返回后端支持的字词级时间戳;Qwen CUDA vLLM 与 CPU Rust 在启用时会自动调用 forced aligner |
| `hotwords` | string | - | 热词,格式:`词1 权重1 词2 权重2` |
| `response_format` | string | `verbose_json` | 输出格式 |
| `prompt` | string | - | 提示文本(保留兼容) |
| `temperature` | float | `0` | 采样温度(保留兼容) |
**音频/视频输入方式:**
- **文件上传**: 使用 `file` 参数上传音频文件或带音轨的视频容器
- **URL / 本地路径读取**: 使用 `audio_address` 参数提供音频/视频 URL 或服务端本地路径,服务将自动读取
- **优先级**: 如果同时提供 `file` 和 `audio_address`,服务会优先使用 `file`,并忽略 `audio_address`
**使用示例:**
```python
# 使用 OpenAI SDK
from openai import OpenAI
client = OpenAI(base_url="http://localhost:8000/v1", api_key="your_api_key")
with open("audio.wav", "rb") as f:
transcript = client.audio.transcriptions.create(
file=f,
response_format="verbose_json" # 获取分段和说话人信息
)
print(transcript.text)
```
```bash
# 使用 curl
curl -X POST "http://localhost:8000/v1/audio/transcriptions" \
-H "Authorization: Bearer your_api_key" \
-F "file=@audio.wav" \
-F "model=qwen3-asr-0.6b" \
-F "response_format=verbose_json" \
-F "enable_speaker_diarization=true" \
-F "enable_speaker_identification=true" \
-F "enable_text_cleanup=true" \
-F "hotwords=Qwen 2.0 ModelScope 1.5"
```
**支持的响应格式:** `json`, `text`, `srt`, `vtt`, `verbose_json`
### 阿里云兼容接口
| 端点 | 方法 | 功能 |
| ------------------------- | --------- | ---------------------- |
| `/stream/v1/asr` | POST | 语音识别(支持长音频) |
| `/stream/v1/asr/models` | GET | 声明条目列表 |
| `/stream/v1/asr/health` | GET | 健康检查 |
| `/ws/v1/asr` | WebSocket | Qwen3-ASR 流式识别 |
| `/ws/v1/asr/qwen` | WebSocket | Qwen3-ASR 流式识别(显式路径) |
| `/ws/v1/asr/funasr` | WebSocket | 已移除;会返回废弃错误并提示切换到 `/ws/v1/asr/qwen` |
**请求参数:**
| 参数 | 类型 | 默认值 | 说明 |
| ------------------------------ | ------ | ------------------ | ------------------------------------- |
| `audio_address` | string | `https://media.cdn.vect.one/podcast_demo.mp4`(文档示例) | 音频/视频 URL、`file://` 或服务端本地路径(可选;若同时上传内容则忽略) |
| `sample_rate` | int | `16000` | 采样率 |
| `enable_speaker_diarization` | bool | `true` | 启用说话人分离 |
| `enable_speaker_identification` | bool | `true` | 说话人分离开启时匹配已注册声纹库 |
| `enable_text_cleanup` | bool | `true` | 启用文本去重、跨段重叠裁剪和口头语清理 |
| `word_timestamps` | bool | `false` | 返回后端支持的字词级时间戳;Qwen CUDA vLLM 与 CPU Rust 在启用时会自动调用 forced aligner |
| `vocabulary_id` | string | - | 热词(格式:`词1 权重1 词2 权重2`) |
**使用示例:**
```bash
# 基本用法
curl -X POST "http://localhost:8000/stream/v1/asr" \
-H "Content-Type: application/octet-stream" \
--data-binary @audio.wav
# 带参数
curl -X POST "http://localhost:8000/stream/v1/asr?enable_speaker_diarization=true&enable_speaker_identification=true&enable_text_cleanup=true&vocabulary_id=Qwen%202.0%20ModelScope%201.5" \
-H "Content-Type: application/octet-stream" \
--data-binary @audio.wav
```
### 会议离线接口
| 端点 | 方法 | 功能 |
| ---- | ---- | ---- |
| `/api/v1/asr/transcriptions` | POST | 创建离线会议识别任务 |
| `/api/v1/asr/transcriptions/{task_id}` | GET | 查询任务状态和结果 |
该接口生产调用只使用 `audio_address`。
```json
{
"audio_address": "https://example.com/media/meeting.mp4",
"config": {
"enable_speaker": true,
"match_speaker_registry": true,
"enable_text_cleanup": true,
"speaker_threshold": 0.6,
"word_timestamps": false,
"hotwords": [
{ "hotword": "通义千问", "weight": 2.0 },
{ "hotword": "ModelScope", "weight": 1.5 }
]
}
}
```
**响应示例:**
```json
{
"task_id": "xxx",
"status": 200,
"message": "SUCCESS",
"result": "说话人1的内容...\n说话人2的内容...",
"duration": 60.5,
"processing_time": 1.234,
"segments": [
{
"text": "今天天气不错。",
"start_time": 0.0,
"end_time": 2.5,
"speaker_id": "说话人1",
"word_tokens": [
{"text": "今天", "start_time": 0.0, "end_time": 0.5},
{"text": "天气", "start_time": 0.5, "end_time": 0.9},
{"text": "不错", "start_time": 0.9, "end_time": 1.3}
]
}
]
}
```
## 说话人分离
基于 CAM++ 模型实现多说话人自动识别:
- **默认开启** - `enable_speaker_diarization=true`
- **自动识别** - 无需预设说话人数量,模型自动检测
- **说话人标记** - 响应中包含 `speaker_id` 字段(如 "说话人1"、"说话人2")
- **智能合并** - 两层合并策略避免孤立短片段:
- 第一层:小于10秒的同说话人片段累积合并
- 第二层:连续片段累积合并至60秒上限
- **字幕支持** - SRT/VTT 格式输出包含说话人标记 `[说话人1] 文本内容`
关闭说话人分离:
```bash
# OpenAI API
-F "enable_speaker_diarization=false"
# 阿里云 API
?enable_speaker_diarization=false
```
## 音频处理
### 智能分段策略
长音频自动分段处理:
1. **VAD 语音检测** - 检测语音边界,过滤静音
2. **贪婪合并** - 累积语音段,确保每段不超过 `MAX_SEGMENT_SEC`(默认60秒)
3. **静音切分** - 语音段间静音超过3秒时强制切分,避免包含过长静音
4. **批处理推理** - 多片段并行处理,GPU 模式下性能提升 2-3 倍
### WebSocket 流式识别限制
**Qwen3-ASR 流式**(使用 `/ws/v1/asr` 或 `/ws/v1/asr/qwen`):
- ✅ 支持多语言实时识别
- ✅ 当前支持 CUDA vLLM 与 CPU Rust 两条流式路径
- ❌ 当前流式路径不返回词级时间戳
### Qwen3 运行时矩阵
| 运行环境 | 后端 | 离线转写 | WebSocket 流式 | 离线词级时间戳 | 流式词级时间戳 | 成熟度 |
|---------|------|---------|----------------|----------------|----------------|--------|
| Linux + NVIDIA GPU | 官方 vLLM 0.20.0 | ✅ | ✅ | ✅ | ❌ | 面向生产 |
| CPU / macOS | QwenASR Rust | ✅ | ✅ | ✅(forced aligner) | ❌ | 推荐本地后端 |
## 支持离线的模型
| 模型 ID | 名称 | 说明 | 特性 |
| -------------------- | ----------------- | ---------------------------------------- | --------- |
| `qwen3-asr-1.7b` | Qwen3-ASR 1.7B | 高性能多语言 ASR;CUDA 使用 vLLM | 离线/实时 |
| `qwen3-asr-0.6b` | Qwen3-ASR 0.6B | 轻量版多语言 ASR;CUDA 使用 vLLM,CPU/macOS 使用 Rust backend | 离线/实时 |
**运行时选择:**
- **显存 >= 32GB**: 选择 `qwen3-asr-1.7b`
- **显存 < 32GB**: 选择 `qwen3-asr-0.6b`
- **无 CUDA**: 选择基于 vendored Rust 的 `qwen3-asr-0.6b`
- **macOS / Apple Silicon**: 无论内存大小多少,默认都选择 `qwen3-asr-0.6b`
- **环境变量覆盖**: 设置 `QWEN3_ASR_MODEL=qwen3-asr-1.7b` 或 `QWEN3_ASR_MODEL=qwen3-asr-0.6b` 可跳过自动选择
启动时会先检测当前运行计划所需模型;如果本地缓存缺失,会自动从 ModelScope 下载。离线部署请提前准备模型缓存。
## 环境变量
推荐直接关心的公开配置:
| 变量 | 默认值 | 说明 |
| ---------------------------------- | ------------ | ----------------------------------------------- |
| `API_KEY` | - | API 认证密钥(可选,未配置时无需认证) |
| `LOG_LEVEL` | `INFO` | 日志级别(DEBUG/INFO/WARNING/ERROR) |
| `MAX_AUDIO_SIZE` | `2048` | 最大音频文件大小(MB,支持单位如 2GB) |
| `ASR_BATCH_SIZE` | `4` | 长音频分段后的 ASR 批处理大小 |
| `MAX_SEGMENT_SEC` | `60` | 音频分段最大时长(秒) |
| `ASR_ENABLE_NEARFIELD_FILTER` | `true` | 启用远场声音过滤 |
| `QWEN3_ASR_MODEL` | 自动选择 | 强制选择 `qwen3-asr-1.7b` 或 `qwen3-asr-0.6b` |
| `QWEN_GPU_MEMORY_UTILIZATION` | `0.9` | vLLM 可保留的 GPU 显存上限;共享显卡时可调低,KV cache 不足时可适当调高 |
| `QWEN_VLLM_ENFORCE_EAGER` | `true` | 强制 vLLM eager 执行以提高兼容性;NVIDIA 性能测试可设为 `false` 允许 CUDA Graph 优化 |
远场过滤调优建议:
- `ASR_NEARFIELD_RMS_THRESHOLD=0.01` 是当前默认值,也是推荐起点
- 嘈杂环境可以适当调高,增强背景语音过滤
- 安静环境如果出现小声说话漏识别,可以适当调低
- 需要观察过滤行为时,可临时设置 `LOG_LEVEL=DEBUG`
后端专项高级配置:
| 变量 | 默认值 | 说明 |
| --- | --- | --- |
| `QWEN_RUST_CPU_WORKERS` | `4` | CPU Rust backend worker 数(Rust ASR / forced align 默认 4 个 runtime) |
| `QWENASR_LIBRARY_PATH` | 自动探测 | 覆盖 vendored Rust 动态库路径 |
## 资源需求
**最小配置(CPU):**
- CPU: 4 核
- 内存: 16GB
- 磁盘: 20GB
**推荐配置(GPU):**
- CPU: 4 核
- 内存: 16GB
- GPU: NVIDIA GPU (16GB+ 显存)
- 磁盘: 20GB
## API 文档
启动服务后访问:
- Swagger UI: `http://localhost:8000/docs`
- ReDoc: `http://localhost:8000/redoc`
## 相关链接
- **部署指南**: [详细文档](./deployment.md)
- **Qwen3-ASR**: [Qwen3-ASR GitHub](https://github.com/QwenLM/Qwen3-ASR)
- **FunASR**: [FunASR GitHub](https://github.com/alibaba-damo-academy/FunASR)
- **QwenASR**: [QwenASR GitHub](https://github.com/huanglizhuo/QwenASR)
## 许可证
本项目采用 MIT 许可证 - 查看 [LICENSE](../LICENSE) 文件了解详情。
## Star 历史
[![Star History Chart](https://api.star-history.com/svg?repos=Quantatirsk/qwen3-asr&type=Date)](https://star-history.com/#Quantatirsk/qwen3-asr&Date)
## 贡献
欢迎提交 Issue 和 Pull Request 来改进项目!

View File

@ -0,0 +1,382 @@
# 实时 ASR WebSocket 处理细节对照
本文只整理当前代码,不修改 WebSocket 或说话人算法。对照对象是:
- 当前独立 Demo:`demo/realtime_asr_optimization_demo`
- 原项目实时接口:`app/api/v1/websocket_asr.py`、`app/services/qwen3_websocket_asr.py`、`app/services/realtime_speaker_clusterer.py`
代码是本文的依据;旧的排查记录或早期说明如果与当前实现冲突,以代码为准。
## 1. 先看整体差异
```mermaid
sequenceDiagram
participant B as 浏览器
participant D as 当前 Demo /ws
participant V as 独立 vLLM HTTP
participant A as 独立辅助服务
B->>D: start(扁平字段)
B->>D: PCM/WAV 二进制帧
D->>D: 20ms RMS 门控、turn 缓冲、静音切段
D->>V: 当前 turn 的累积 WAV(partial/final)
V-->>D: 文本
D-->>B: sentences + display_state
D->>A: final turn + session_id
A-->>D: CAM++ embedding / 在线聚类标签
D-->>B: 同 sentence_id 的 speaker 更新 + 新 display_state
B->>D: eof 或 stop
D->>D: 等待音频队列和 speaker 队列清空
D-->>B: end
```
```mermaid
sequenceDiagram
participant B as 浏览器
participant R as 原项目 FastAPI 路由
participant Q as Qwen3ASRService
participant E as Qwen3ASREngine
participant S as RealtimeSpeakerClusterer
B->>R: /ws/v1/asr 或 /ws/v1/asr/qwen
R->>Q: handle_connection
B->>Q: start.payload(嵌套字段)
Q-->>B: voice_id、start
B->>Q: PCM/WAV 二进制帧
Q->>Q: 转采样、VAD、pre-roll、partial
Q->>E: 原生流式或当前窗口重转写
E-->>Q: partial 文本
Q-->>B: sentence_type=0
Q->>E: 当前 turn 全量 final 重转写
E-->>Q: final 文本
Q-->>B: sentence_type=1、speaker_id=-1
Q->>S: speaker_job_queue(异步)
S-->>Q: 聚类/注册库匹配
Q-->>B: 相同 sentence_id 的 speaker 回写
B->>Q: stop
Q-->>B: end(stop 内部会再次 final/recluster)
```
核心设计差别是:当前 Demo 把 ASR 和声纹都放在独立 HTTP 服务后面,WebSocket 只做编排;原项目把 Qwen ASR 引擎、VAD、在线声纹聚类放在同一个服务进程里,但 speaker 归属仍是 final 之后异步补回。
## 2. 当前独立 Demo 的处理链路
### 2.1 进程和启动关系
`server.py` 启动一个 aiohttp HTTP/WebSocket 进程,并在生命周期中创建两个 HTTP 客户端:
| 组件 | 默认地址 | 职责 |
| --- | --- | --- |
| `VLLMTranscriptionService` | `http://127.0.0.1:9950/v1` | 只调用 `/audio/transcriptions`,partial 和 final 都是累积窗口 HTTP 请求 |
| `AuxiliaryModelService` | `http://127.0.0.1:8010` | 调用 `/health`、`/v1/speaker/resolve`、`/v1/speaker/reset` |
| `RealtimeSession` | `server.py:/ws` | 接收音频、VAD 切 turn、调用两个服务、维护展示状态 |
辅助服务在 `demo/scripts/auxiliary_server.py` 中预加载 VAD 与 CAM++ `speaker_verification`。每次 resolve 只上传一个已结束 turn,服务以 `session_id` 保存在线聚类中心;这不是把整段会议音频重新上传。
### 2.2 WebSocket 输入和输出
客户端首条消息是扁平结构:
```json
{
"type": "start",
"source": "mic",
"model_service_url": "http://127.0.0.1:9950/v1",
"model": "Qwen/Qwen3-ASR-0.6B",
"speaker_diarization": 1,
"sentence_strategy": 0,
"partial_interval_ms": 1200,
"max_segment_sec": 12,
"display_merge": true
}
```
`start` 成功后,客户端发送 16kHz、单声道、PCM16 二进制帧。文件模式只接收 `.pcm` 和 `.wav`;WAV 的 RIFF/fmt/data chunk 在 WebSocket 服务端增量剥离,并且要求 16kHz、单声道、PCM16。
控制消息:
| 消息 | 行为 |
| --- | --- |
| `eof` | 输入生产者结束;服务端把 `EOF` 放入音频队列,完成尾部 turn 和 speaker 队列后发送 `end` |
| `stop` | 与 `eof` 走同一排空流程,同时记录 `input_stopped`,用于 WAV 不完整时的校验差异 |
| `abort` | 取消音频和 speaker worker,直接结束,不保证当前 turn 有 final |
服务端消息的实际顺序通常是:
```text
start
-> sentences(partial,可能多次)
-> display_state(每次状态改变一份快照)
-> sentences(final,speaker 尚未确认)
-> display_state(pending)
-> display_state(processing/confirmed 或失败原因)
-> draining
-> end
```
`sentences` 是兼容性事件;`display_state` 是当前页面的主要渲染数据。`display_state.revision` 单调递增,包含:
- `raw_segments`:按 `start_time`、`sentence_id` 排序的原始片段
- `display_blocks`:按相邻且可信的说话人合并后的展示块
- `metrics`:音频字节数、输入帧数、partial 数量、partial 修订次数和耗时
### 2.3 音频、VAD 和 turn 边界
`RealtimeSession.process_audio()` 的边界是本 Demo 最重要的状态机:
1. 二进制数据进入 `audio_queue`,再拆成 640 bytes 的 PCM 帧,即 20ms。
2. 每帧用 RMS 阈值 `450` 判断有声/静音;这是 WebSocket 层的轻量门控,不是辅助服务的整段 VAD pipeline。
3. 未进入说话状态时,保留最近 6400 bytes(约 200ms)`pre_roll`。
4. 第一帧有声时,将 pre-roll 加到新 `segment_audio`,设置 `segment_start_ms`。
5. 进入说话状态后,所有帧追加到当前 turn;有声帧累计 `voiced_ms`,静音帧累计 `silence_ms`。
6. `sentence_strategy=0` 默认约 800ms 静音提交;`sentence_strategy=1` 使用约 1400ms 静音提交。
7. 达到 `partial_interval_ms`(默认 1200ms)且尚未达到静音阈值时,调用一次 vLLM partial。
8. 达到 `max_segment_sec`(默认 12s)时按 `max_duration` 提交。
当前 Demo 不把切段尾部静音送给 ASR/声纹:提交前按 `silence_ms` 从 `segment_audio` 尾部删除。下一个 turn 没有原项目那样的 `carry_audio`,而是从后续有声帧重新开始;因此两段之间的静音会形成时间间隔,不会自动带入下一段。
### 2.4 ASR 结果和同句覆盖
`_emit_transcription()` 始终使用当前 `segment_id`:
- `sentence_type=0`:partial,写入/覆盖同一个 `sentence_id`
- `sentence_type=1`:final,仍写入同一个 `sentence_id`,并加入 speaker 队列
`SegmentAssembler.apply_sentence()` 先按 `sentence_id` 找旧记录再覆盖;如果已经是 final,后来的 partial 不会回滚 final。final 没有文本时,会移除该片段,避免遗留一个永久 pending 的 partial。
### 2.5 声纹异步链路
`_commit_segment()` 只负责把 `SpeakerJob` 放入 `speaker_queue`,不会等待 CAM++。`process_speakers()` 是单 worker,按 turn 入队顺序串行调用辅助服务:
1. `voiced_ms < 800ms`:不提取声纹,写入 `insufficient_audio`,保持 `speaker_id=-1`。
2. 辅助服务缺失或异常:写入 `service_unavailable`/`service_error`,ASR 继续输出。
3. 辅助服务返回 embedding/聚类结果:回写同一 `sentence_id`。
4. `SegmentAssembler.apply_speaker_update()` 只接受可信身份:`speaker_evidence` 必须是 `fresh` 或 `confirmed`,置信度至少 `0.6`,并拒绝 `short_attach`、`embedding_attach`。
5. 不可信结果被归一化为 `speaker_id=-1`、`speaker_name=""`,但保留状态和原因供诊断。
辅助服务的在线聚类是简单的 session 级中心匹配:首次 embedding 新建 `speaker_id`,之后与已有中心的余弦相似度达到阈值就更新中心并复用 ID。`/v1/speaker/reset` 在 WebSocket 结束时清理该 session。
### 2.6 展示块如何合并
`SegmentAssembler.display_blocks(merge_adjacent=True)` 先按时间排序原始片段,然后遵循:
- pending/unknown 片段始终以自己的 `sentence_id` 作为身份键,独立成块;
- 只有相邻且可信的片段,且 `user_id`/`registry_speaker_id`/`speaker_id` 身份键相同,才合并;
- 合并只拼接文本、扩大结束时间并追加 `segment_ids`,原始片段仍保留在 `raw_segments`。
当前页面收到 `display_state` 后会整块重建结果区,按 `block_id` 渲染;收到 `sentences` 时如果服务端声明支持 `display_state`,页面不会再次追加,避免同一片段重复显示。旧页面或绕过 `display_state` 的客户端不具备这个保护。
## 3. 原项目 WebSocket 的处理链路
### 3.1 路由和状态
`app/api/v1/websocket_asr.py` 当前实际路由:
- `/ws/v1/asr`
- `/ws/v1/asr/qwen`
- `/ws/v1/asr/funasr` 已废弃,接受后发送 `FUNASR_REALTIME_REMOVED` 并以 1008 关闭
每个连接进入 `Qwen3ASRService.handle_connection()`。`ConnectionContext.state` 为 `READY -> STARTED -> STREAMING`;会话还可以通过 `session_id` 在 TTL 内断线恢复。恢复的是完整上下文,包括已确认片段、speaker history、时间线和待处理状态,不只是一个 WebSocket ID。
### 3.2 start 参数
原项目的 `start` 使用 `payload` 嵌套对象。常用字段包括:
| 类别 | 字段 |
| --- | --- |
| 音频 | `format`、`sample_rate`、`language`、`context`、`enable_inverse_text_normalization` |
| partial | `min_partial_sec`、`partial_window_sec`、`partial_holdback_chars`、`unfixed_token_num`、`enable_native_partial_stream` |
| 切段 | `silence_duration_ms`(默认 800)、`pre_roll_ms`(默认 240)、`max_sentence_count`(默认 8)、`enable_realtime_vad_split`、`max_segment_sec` |
| speaker | `enable_speaker`(默认 true)、`match_speaker_registry`、`speaker_threshold` |
| 稳定性 | `force_stable_segment_sec`、`force_stable_min_chars`、`soft_limit_sec`、`hard_limit_sec` |
服务端先发送 `voice_id`,再发送 `start`。`voice_id` 与会话 ID 通常相同;如果客户端传入固定 `payload.session_id`,断线重连时可以复用上下文。
### 3.3 音频转换、pre-roll 和 partial
`_convert_audio()` 接收 PCM/WAV,转为 float32;多声道下混为单声道,非 16kHz 用 scipy 重采样。每个二进制消息都在服务端转换后立即参与 VAD。
原项目的 `ConnectionContext` 同时维护:
- `pre_roll_audio`:未开始说话前的前滚音频,默认约 240ms;
- `segment_audio_buffer`:从当前 turn 开始到提交前的完整音频;
- `stream_window_buffer`:最近窗口,用于 partial 或 native stream 失败时回退;
- `realtime_stream_state`:只服务低延迟 partial,不决定 final;
- `silence_samples`、`sentence_active`、`total_samples`:VAD 状态和当前 turn 长度。
有声输入到达时,`_start_turn()` 把 pre-roll 与当前音频拼接;后续由 `_append_turn_audio()` 追加。达到最短窗口后,服务端可走原生 Qwen partial,或对当前窗口/当前 turn 重转写;partial 会经过清理、去重、与上一段重叠裁剪后发送。
提交触发条件不只有静音:
- 识别到足够完整的标点句,且时长/字数达到稳定门槛;
- 句子数达到 `max_sentence_count`;
- 静音样本达到 `silence_duration_ms`;
- 达到硬时长限制;开启实时 VAD split 时会尝试找一个完成的分割点。
### 3.4 final、carry 和时间线
`_commit_retranscribe_turn()` 对当前完整 turn 做一次 final 重转写,生成 `confirmed_segments` 元素:
```text
index / text / language / reason
duration_ms / start_ms / end_ms
sentence_type=1 / speaker_id=-1 / speaker_pending=true
```
final 事件先发送,speaker 之后再补。默认 final 的 `start_ms` 来自 `ctx.timeline_cursor_ms`;提交后时间线前移到 `segment_end_ms`。
开启实时 VAD split 时,提交可能得到 `finalized_audio + carry_audio`:前半段定稿,后半段留在下一个逻辑 turn 中,且会重新初始化 stream 状态。这个 carry 是原项目与当前 Demo 的一个实质差异,也是跨说话人边界时必须重点观察的音频来源。
### 3.5 原项目 speaker worker 和聚类
原项目 final 后把 job 放入 `ctx.speaker_job_queue`,由 `_speaker_worker_loop()` 串行消费。`_resolve_segment_speaker()` 调用 `RealtimeSpeakerClusterer.resolve_segment_speaker()`,再把结果写入指定 `segment_index`,通过相同索引发送一条新的 `sentences`。
`RealtimeSpeakerClusterer` 当前行为:
- 1.6s 以下且上一条有命名身份:使用 `short_attach` 直接沿用上一条;
- 正常 turn:按 1.5s 窗口、0.75s 步长提取 CAM++ embedding,匹配已有记录或新建 generic speaker;
- 4s 以上若 chunk 明显混合:返回 `mixed_segment`,保持未知;
- 开启注册库匹配且时长至少 2.4s:在独立注册 embedding 空间匹配实名;
- timeline 平滑时,短于 0.7s 的范围会并给相邻说话人;
- 实时记录达到至少 5 条且队列积压不超过 1 条时,可能对最近 12 条 pending 片段重新聚类;
- `stop` 时还会对历史记录做一次最终 recluster。
此外,`Qwen3ASRService` 自身还有两类“最近说话人继承”:generic speaker 新 turn 时长至少 8s 才允许 `recent_inherit`,实名 speaker 至少 4.5s 才允许 `recent_named_inherit`。这些继承都发生在 embedding 结果之后,不能与 `short_attach` 混为一谈。
### 3.6 stop 和 end
客户端只发送 `{"type":"stop"}`。服务端会:
1. 对仍 active 的 `segment_audio_buffer` 做 `reason=final` 的 final 提交;
2. 立即对现有 `speaker_records` 做最终 recluster,并发送可能的 speaker 更新;
3. 汇总 `confirmed_segments`,发送 `end(final=1)`。
这里与当前 Demo 不同:原项目的 `_stop()` 没有显式等待 `speaker_job_queue.join()`。如果 stop 到达时 speaker worker 仍在处理,最终 recluster/end 可能先于某个异步 speaker 回写;断开清理还会停止 worker。客户端必须把同 `sentence_id` 的后续 speaker 事件当作可迟到更新,而不能认为 `end` 之后绝不会再有归属变化。
## 4. 两套消息契约对照
| 维度 | 当前独立 Demo | 原项目 |
| --- | --- | --- |
| WebSocket | aiohttp `/ws` | FastAPI `/ws/v1/asr`、`/qwen` |
| start | 扁平字段 | `payload` 嵌套字段 |
| ASR | 外部 vLLM HTTP 累积窗口 | 进程内 Qwen engine,原生 stream 或重转写回退 |
| VAD | 20ms PCM RMS 门控 | float32 音频门控,支持实时 VAD split 辅助切分 |
| pre-roll | 固定约 200ms | `pre_roll_ms` 默认约 240ms |
| final 音频 | 删除提交尾部静音,不保留 carry | 可有 `carry_audio` 并带入后续 turn |
| speaker | 外部 CAM++/在线中心服务,单 worker | 进程内 CAM++ chunk 聚类、注册库、重聚类,单 worker |
| 未确认 speaker | `speaker_evidence=pending`,展示独立未知块 | `speaker_id=-1` 或 `speaker_pending=true`,客户端需自行暂存 |
| 文本更新键 | `sentence_id` | `sentence_id` 对外,内部 `segment_index` |
| 展示快照 | `display_state.revision`,服务端生成 `display_blocks` | 没有同等的服务端展示块协议,客户端按 sentence upsert |
| end 屏障 | `EOF -> audio worker -> speaker EOF -> end` | `_stop()` final/recluster/end,不等待 speaker 队列清空 |
## 5. “上一人的最后一句进入下一人气泡”的定位框架
当前先不修改,定位时要把“片段本身错了”和“片段正确但展示合错了”分开。
### 5.1 当前独立 Demo 的可能路径
1. **同一个 turn**:两人换话之间没有达到 800ms(或段落模式 1400ms)静音,RMS VAD 不切段。此时 vLLM 收到的是混合 turn,前端只有一个 `sentence_id`,不是气泡合并问题。
2. **声纹误归属**:A 的 final 先是 pending,随后辅助服务把 A 误匹配到 B 的 cluster。下一次 `display_state` 中,A、B 两个相邻片段拥有同一可信身份,`display_blocks()` 会把两者拼成一个 block。
3. **异步回写改变了合并条件**:A 的 speaker 结果可能在 B 的 final 之后才到达。服务端按时间排序重建快照,所以视觉上是 A 的文字“后来进入”B 的气泡;实际是 A 的旧片段身份被补齐后触发了相邻合并。
4. **旧客户端渲染路径**:当前页面在 `display_state_supported=true` 时忽略 `sentences`,但旧页面若逐条 append `sentences`,可能把同一 `sentence_id` 的 final/speaker 更新当成新气泡,或把 pending 文本追加到上一气泡。必须确认浏览器加载的 `app.js` 版本和服务端返回的 `display_state`。
5. **session 污染**:辅助服务按 `session_id` 保存聚类中心。若 reset 没有执行、多个连接错误复用同一个 session ID,上一场会话的 cluster 可能影响新会话;正常一次连接内 A/B 共用中心是设计行为,不是跨人合并的充分证据。
当前 Demo 的 worker 是串行的,`emit()` 有发送锁,因而“并发返回顺序打乱”不是首要嫌疑;首要证据应是 `raw_segments` 的 `sentence_id/start_time/speaker_id/speaker_strategy` 是否正确,以及 `display_blocks.segment_ids` 是否把两个片段合到一起。
### 5.2 原项目的可能路径
1. **静音边界不足**:默认 800ms 静音才提交;换话前的短停顿会让 A 尾部和 B 开头留在同一个 `segment_audio_buffer`。
2. **carry 音频污染**:启用实时 VAD split 时,分割点之后的 `carry_audio` 会成为下一个 turn 的开头。若 split 点落在 A 尾音或 B 起音中间,下一段声纹和 ASR 都会携带前一人尾部。
3. **短句沿用上一身份**:B 的新段小于 1.6s 且历史有命名 speaker 时,`short_attach` 会直接复用上一条;chunk 提取为空时的 `embedding_attach` 也可能复用上一条。这是代码中最直接的“上一人污染下一段”路径。
4. **最近身份继承**:较长的新段在匹配失败时可能触发 `recent_named_inherit`(至少 4.5s)或 `recent_inherit`(至少 8s),因此不能只看最终 `speaker_id`,还要记录 `speaker_strategy`。
5. **重聚类改写历史**:实时重聚类和 stop 最终重聚类都可能改写已有 segment 的 speaker。客户端如果按到达顺序追加,而不是按 `sentence_id` upsert,就会看到旧气泡和新气泡互相覆盖或合并。
6. **stop 竞态**:原项目 stop 不等待 speaker job 队列清空;end 可能先发,随后连接清理还会取消 worker。最后一个人的 speaker 归属可能缺失、迟到或停留在旧标签,前端若把 end 当成不可变快照会放大问题。
### 5.3 需要同时保存的证据
对同一段测试音频,至少保存以下三层结果:
```text
音频层:每个 turn 的 start/end、有效有声时长、是否包含 carry/pre-roll
识别层:sentence_id/index、sentence_type、文本、speaker_strategy、置信度
展示层:display_state.revision、raw_segments、display_blocks.segment_ids
```
判定规则:
- `raw_segments` 已经只有一个片段:先查 VAD/切段边界;
- `raw_segments` 有 A、B 两段且 speaker ID 相同:查声纹误匹配/继承/重聚类;
- `raw_segments` 的 ID 不同但 `display_blocks.segment_ids` 合并:查展示合并键;
- `display_blocks` 正确但页面仍显示一只气泡:查浏览器脚本版本、是否绕过 `display_state`、是否按 `sentence_id` upsert。
## 6. 后续细节优化的优先级(本轮不实施)
### P0:先证明边界和身份是否正确
- 记录每个 final 的实际音频起止、有效有声毫秒、pre-roll/carry 长度。
- 记录 speaker resolve 请求和返回的 `session_id`、策略、置信度、cluster ID。
- 前端临时展示 `block.segment_ids`,确认“合并”到底是两个片段还是一个片段。
- 用固定 A-静音-B 音频,比较 200ms、500ms、800ms、1400ms 停顿。
### P1:降低错误身份传播
- 对 `short_attach`、`embedding_attach`、recent inherit 单独统计,不要只统计 speaker_id。
- 对跨边界的短 turn 保持 pending,等到有独立 embedding 或后续重聚类再确认。
- 明确实时重聚类和最终重聚类的可修改范围,客户端统一按 ID 幂等更新。
- 为 stop 增加“最后一个 speaker job 已完成”的可观察状态。
### P2:改善展示稳定性
- 展示层只把可信且相邻的片段合并,保留 segment_ids 和 revision。
- 对 speaker 更新做局部重绘或整快照重绘,但不要把同一个 sentence 当成新消息追加。
- unknown/pending 使用独立块,不把诊断文本放在说话人名称中。
## 7. 建议的验收用例
| 用例 | 观察点 | 通过标准 |
| --- | --- | --- |
| A 说 3s,停 1s,B 说 3s | 两套服务的 raw segment | 至少两个不同 sentence_id,时间不重叠 |
| A 说 3s,停 300ms,B 说 3s | VAD 边界 | 明确记录为同段或分段,不能只看气泡颜色判断 |
| A 说 3s,B 只说 0.8s | 短 turn speaker 策略 | 原项目应能观察 `short_attach`;Demo 应保持 pending 或独立结果 |
| A/B 各说多段,speaker 服务延迟 2s | 异步回写 | 文本不重复,更新按 sentence_id 定位,顺序按时间恢复 |
| speaker 服务不可用 | 降级 | ASR 仍有 final,speaker 为未知并有明确 reason |
| stop 紧跟最后一帧 | 收尾屏障 | Demo 的 end 在 speaker 队列完成后发送;原项目记录可能迟到的 speaker 更新 |
| 断线后同 session_id 重连 | 会话隔离 | 原项目按 TTL 恢复;Demo 新连接不会复用旧 speaker center |
## 8. 源码索引
### 当前独立 Demo
- `demo/realtime_asr_optimization_demo/server.py`
- `RealtimeSession.__init__`:会话参数、队列和 VAD 状态
- `emit_state`:`raw_segments/display_blocks/revision`
- `_emit_transcription`:partial/final 写入同一 `sentence_id`
- `_commit_segment`、`process_audio`:VAD、切段和 speaker job 入队
- `_resolve_speaker`、`process_speakers`:异步声纹回写
- `websocket_handler`:start、二进制帧、eof/stop/abort、end
- `demo/realtime_asr_optimization_demo/speaker_assembler.py`
- `apply_sentence`、`apply_speaker_update`、`display_blocks`
- `demo/realtime_asr_optimization_demo/model_service.py`
- 独立 vLLM OpenAI-compatible HTTP 适配
- `demo/realtime_asr_optimization_demo/auxiliary_service.py`
- `/health`、`/v1/speaker/resolve`、`/v1/speaker/reset` 客户端适配
- `demo/scripts/auxiliary_server.py`
- VAD/CAM++ 预加载、embedding 提取、session 级在线聚类
### 原项目
- `app/api/v1/websocket_asr.py`
- `/ws/v1/asr`、`/ws/v1/asr/qwen`、废弃 `/funasr`
- `app/services/qwen3_websocket_asr.py`
- `ConnectionContext`:音频、partial、confirmed segments、speaker 队列
- `handle_connection`:WebSocket 状态机和消息协议
- `_commit_retranscribe_turn`:final、carry、timeline、speaker job
- `_speaker_worker_loop`、`_resolve_and_emit_segment_speaker`:异步 speaker 回写
- `_inherit_recent_*`、`_maybe_recluster_recent_segments`:身份传播和重聚类
- `_stop`:最终提交、recluster、end
- `app/services/realtime_speaker_clusterer.py`
- chunk embedding、短段沿用、已有 speaker 匹配、混合段、timeline 平滑
- `docs/realtime_meeting_websocket.md`
- 对外协议示例;其中 partial/final/speaker 回写必须按 `sentence_id` 幂等处理
本轮只新增本文档,没有修改上述实现。后续修复应先用第 5 节的三层证据确定问题属于切段、声纹还是展示层,再决定改哪一层。

View File

@ -0,0 +1,66 @@
# 熵减执行清单
本文档记录本轮已执行的熵减工作。目标是移除不可达路径、兼容占位、隐藏 fallback、重复请求流水线和无用依赖。
## 范围
- 主范围:`app/`、根运行配置、公开运行文档。
- 不处理:`vendor/qwenasr` 内部实现、仅 benchmark 使用且不阻塞主链路的代码。
- 原则:开发中项目不保留废弃接口、旧字段、兼容层或 fallback 逻辑。
## P0
- [x] 移除 Qwen3 `transformers` 后端残留路径。
- 删除后端选择里的 `"transformers"` fallback。
- 删除仅服务该路径的 batch/segment 转换死代码。
- 不支持的设备显式失败。
- [x] 启动预加载改为 fail-fast。
- 模型完整性检查或预加载失败时停止 worker。
- 不再静默降级到首次请求加载。
- [x] 移除被忽略的离线模型兼容参数。
- 删除 REST `model_id` 兼容处理。
- 删除 OpenAI transcription `model` 兼容处理。
- 运行时模型选择统一由部署计划和 `QWEN3_ASR_MODEL` 控制。
## P1
- [x] 抽出共享离线转写服务。
- API 层只处理协议输入和响应格式。
- 音频准备、`OfflineASRRequest`、runtime 调用和清理边界集中到服务层。
- [x] 合并音频字节处理逻辑。
- `process_from_request` 和 `process_upload_file` 共享私有 byte 处理 helper。
## P2
- [x] 将 `.tsscale` sidecar 隐式耦合改为显式结构化元数据。
- [x] 拆分 WebSocket 路由和 Qwen3 协议服务。
- Qwen3 websocket 状态机移出 API route。
- route 模块仅保留端点注册和 service delegation。
- [x] 删除首轮发现的无用 helper。
- [x] 审计并移除无用直接依赖。
- 根环境和 CPU 环境移除直接依赖 `pydub`、`httpx`。
- 保留 ModelScope/FunASR 动态 runtime 依赖。
- [x] 继续压缩阿里协议 WebSocket service。
- 删除大段注释、未使用状态、未使用参数和死函数。
- 合并重复响应构造。
- 删除重复音频转换。
## 验证
- [x] `uv run python -m py_compile $(find app -name '*.py' -not -path '*/__pycache__/*') start.py`
- [x] `uvx pyright`
- [x] 变更模块 import smoke check。
- [x] 手工 API smoke plan 已记录:
- `/stream/v1/asr`
- `/v1/audio/transcriptions`
- `/ws/v1/asr/funasr`
- `/ws/v1/asr/qwen`
## 手工 Smoke Plan
启动服务并准备模型后执行:
1. `POST /stream/v1/asr`,使用小 WAV request body,确认返回 `result`、`segments`、`duration`、`processing_time`。
2. `POST /v1/audio/transcriptions`,使用 multipart `file` 和 `response_format=verbose_json`,确认返回 OpenAI 风格 `text` 和 `segments`。
3. 连接 `/ws/v1/asr/funasr`,发送阿里兼容 start/audio/stop 消息,确认 sentence 事件仍正常返回。
4. 连接 `/ws/v1/asr/qwen`,发送 start/audio/stop 消息,确认 partial/final 事件仍正常返回。

View File

@ -0,0 +1,457 @@
# x86 Rust Align Optimization Plan
## 当前结论
- 当前工作区的 vendored Rust backend 已经收敛到更接近 upstream `huanglizhuo/QwenASR` 的 Linux/x86_64 路径:
- `release`
- `RUSTFLAGS="-C target-cpu=native"`
- `BLAS/OpenBLAS`
- x86_64 默认 `BF16` decode
- 保留的有意偏离只有两类:
- `SharedQwenModel` / shared model cache
- 中性的 `ffi` feature(`macos-ffi` 仅作为兼容别名保留)
## upstream 参考
- upstream repo: `https://github.com/huanglizhuo/QwenASR`
- inspected commit: `4e85a19b05f034e106a345d279c68f50df718ab8`
## 本机环境
- CPU: `Intel Core i5-13600KF`
- visible CPUs: `14`
- memory: user reported `G.SKILL DDR5-6400`
## 已验证 benchmark
音频:
- `/opt/qwen3-asr/temp/test_assets/podcast_demo_2min_16k.wav`
- duration: `120s`
### decode 路径对照(runtime concurrency = 4)
`INT8 decode`
- total: `173.42s`
- asr: `54.69s`
- align: `118.74s`
- rtf: `1.4452`
来源:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_int8.json`
`BF16 decode`
- total: `125.87s`
- asr: `46.40s`
- align: `79.47s`
- rtf: `1.0489`
来源:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_bf16.json`
结论:
- 在当前这台 x86_64 机器上,`BF16 decode` 明显优于 `INT8 decode`
- 因此 x86_64 默认 decode 路径应保持 `BF16`
### 收敛后的默认路径 benchmark
来源:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_default_after_converge.json`
结果:
- `runtime concurrency = 4`
- total: `131.77s`
- asr: `51.54s`
- align: `80.23s`
- rtf: `1.0981`
- `runtime concurrency = 14`
- total: `128.25s`
- asr: `57.12s`
- align: `71.13s`
- rtf: `1.0688`
结论:
- 当前主瓶颈仍然在 `align`
- `align_sec` 明显大于或接近 `asr_sec`
- `runtime concurrency` 的最优值并不稳定,说明问题不在单纯线程数,而在具体阶段的访问模式和 kernel 行为
## 为什么不继续默认走 INT8
- upstream README 对 Linux/x86_64 的主路径描述是 `BLAS + AVX2/FMA`
- 当前本机实测中,`INT8 decode` 明显慢于 `BF16 decode`
- 说明这台机器上的主瓶颈不只是权重带宽,更多是:
- x86_64 上 INT8 kernel 的有效带宽利用率
- cache / 数据布局
- 实现成熟度差异
## 当前仍需保留的偏离
### 1. Shared model cache
文件:
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/context.rs`
目的:
- 多 runtime / 多 worker 场景下复用只读模型权重
- 避免每个 runtime 重复 mmap / 持有整套权重
### 2. ffi feature
文件:
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/Cargo.toml`
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/lib.rs`
- `/opt/qwen3-asr/Dockerfile.cpu`
目的:
- 让 Linux CPU 集成不再依赖命名不准确的 `macos-ffi`
- 同时保留兼容别名,避免已有脚本立即失效
## align 热点拆解计划
### Phase 1: 阶段级 profiling
状态:已完成
目标:
- 先确认 `align` 的主耗时究竟在哪一段
位置:
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/align.rs`
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/decoder.rs`
需要拆出的阶段:
- `mel_spectrogram`
- `encoder.forward`
- `input_embeds build`
- `decoder_prefill_logits`
- `timestamp argmax extract`
- `fix_timestamps`
验收:
- 2 分钟样本上输出稳定的阶段级耗时表
实际结果:
- 产物:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_align_profile.json`
- `/opt/qwen3-asr/temp/test_logs/qwen_rust_runtime_concurrency_2min_align_profile.stderr`
- 结论:
- `align` 的主热点明确落在 `decoder_prefill_logits`
- `final rms_norm` 和 `lm_head projection` 不是主矛盾
### Phase 2: decoder_prefill_logits 内部分解
状态:已完成
如果 `decoder_prefill_logits` 是主热点,则继续拆分:
- `decoder_prefill`
- final `rms_norm`
- `lm_head projection`
目的:
- 判断到底是 decoder prefill 慢,还是最后分类头 projection 慢
实际结果:
- 产物:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_align_breakdown.json`
- `/opt/qwen3-asr/temp/test_logs/qwen_rust_runtime_concurrency_2min_align_breakdown.stderr`
- 结论:
- 真正的大头是 `decoder_prefill`
- 长段(`seq_len=1801`)时,`attention_ms` 占 `decoder_prefill` 的绝大部分
- 关键样本:
- `decoder_prefill total_ms=79280.78`
- `attention_ms=73769.91`
- `qkv_ms=1336.33`
- `gate_up_ms=1866.68`
- `down_proj_ms=955.94`
### Phase 2.5: 失败尝试记录
状态:已完成并回退
尝试:
- 针对 x86_64 预先物化 prefill 用 F32 权重,避免每次 `align` 反复做 `BF16 -> F32`
结果:
- 长段 `decoder_prefill` 没有稳定收益,反而出现回归
- 这条路径已经回退,不保留在主线代码里
结论:
- 当前瓶颈不是简单的 BF16 权重转换
- 更直接的问题是 multi-token causal attention 的算法路径
### Phase 3: 对热点段做针对性优化
状态:第一轮已完成
根据 profiling 结果,按优先级选一个方向:
1. 如果热点在 `input_embeds build`
- 复用固定 prefix/suffix embeddings
- 减少逐 token 小块 copy
- 降低每段 align 的重复构造开销
2. 如果热点在 `decoder_prefill`
- 检查 `BF16 matvec / attention / swiglu` 的实际热点
- 优化并行粒度或数据布局
3. 如果热点在 `lm_head projection`
- 优先优化 `BF16` classify head 路径
- 避免无价值的整块 materialize
- 但只有在 profiling 证明是主热点后才动
4. 如果热点在后处理
- 精简 `fix_timestamps`
- 减少 `Vec/String` 分配
### Phase 3 实施结果
本轮实际落地的是:
- 文件:
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/kernels/mod.rs`
- 改动:
- 对 BLAS multi-token causal attention 增加长序列专用 batched 路径
- 仅在 `seq_q >= 256` 时启用
- 从“每个 head、每一行 2 次小 GEMM”改为:
- 每个 head 1 次 `Q @ K^T`
- 行级 causal softmax
- 每个 head 1 次 `softmax @ V`
### Phase 3 回归结果
来源:
- 优化后 benchmark:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_attention_opt.json`
- 优化后 profiling:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_attention_opt_profile.json`
- `/opt/qwen3-asr/temp/test_logs/qwen_rust_runtime_concurrency_2min_attention_opt_profile.stderr`
关键对比:
- 之前默认路径(2 分钟样本,基线文件):
- `total=128.25s`
- `asr=57.12s`
- `align=71.13s`
- 优化后:
- `total=76.46s`
- `asr=45.32s`
- `align=31.14s`
- `rtf=0.6372`
attention 热点变化:
- 长段 `seq_len=1801`
- 优化前:
- `decoder_prefill total_ms=79280.78`
- `attention_ms=73769.91`
- 优化后:
- `decoder_prefill total_ms=16625.88`
- `attention_ms=9981.88`
结论:
- 当前这台 i5 上,`align` 的主矛盾已经从“attention 明显失控”收敛到了“attention 仍是第一热点,但已降到可接受量级”
- 这一轮优化是有效的,应该保留
## Phase 4: FFN 路径继续收敛
状态:已完成一轮,并保留有效部分
本轮动作:
- 文件:
- `/opt/qwen3-asr/vendor/qwenasr/crates/qwen-asr/src/decoder.rs`
- 改动:
- 仅对 x86_64 `BF16 prefill` 的 FFN 路径物化共享 F32 权重
- 范围只包括:
- `gate_up_fused`
- `down_weight`
- `QKV` 和 `O-proj` 暂不纳入
回归数据:
- 不带 profiling:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_ffn_opt.json`
- `total=58.61s`
- `asr=29.50s`
- `align=29.11s`
- 带 profiling:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_ffn_opt_profile.json`
- `total=62.07s`
- `asr=32.15s`
- `align=29.92s`
与上一轮 attention-only profiling 对比:
- attention-only:
- `total=64.58s`
- `align=30.86s`
- FFN-opt:
- `total=62.07s`
- `align=29.92s`
长段热点对比(`seq_len=1801`):
- FFN-opt 之前:
- `attention_ms=9981.88`
- `gate_up_ms=2080.10`
- `down_proj_ms=1183.64`
- FFN-opt 之后:
- `attention_ms=9660.15`
- `gate_up_ms=2210.39`
- `down_proj_ms=1166.96`
结论:
- FFN 这轮不是“大收益”,但 `align_sec` 仍然有小幅下降
- 收益不像 attention 优化那样压倒性,更像是小幅收敛
- 当前可以保留,但不值得继续在同一方向上扩大复杂度
## Phase 5: QKV / O-proj 试验与回退
状态:已完成并回退
尝试:
- 在 prefill 中进一步物化 `QKV` 和 `O-proj` 的共享 F32 权重
回归数据:
- 试验版本:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_qkv_opt.json`
- `total=63.12s`
- `align=30.07s`
结论:
- 相比 FFN-opt 版本,没有形成净收益
- 因此这条路径已回退,不保留在主线代码里
## Phase 6: attention 继续细化的两次试验
状态:已完成并回退
### 试验 A:query-block batched attention
尝试:
- 将长序列 batched causal attention 从“整段一次性 `QK^T` / `SV`”改成按 query block 分块执行
结果:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_attention_block_opt.json`
- `total=79.18s`
- `align=33.40s`
结论:
- 这条路径在当前 i5 + OpenBLAS 组合下没有收益
- 增加 GEMM 次数带来的额外调度开销,超过了小块缓存收益
- 已回退
### 试验 B:提高 batched 切换阈值到 512
尝试:
- 只改 `BATCHED_CAUSAL_ATTENTION_THRESHOLD`
- 让中等长度序列继续走 row-wise 路径,只把更长的序列交给 batched 路径
结果:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_attention_threshold512.json`
- `total=74.49s`
- `align=30.42s`
结论:
- 相比当前主线最优版本也没有收益
- 说明当前阈值 `256` 不是主要问题
- 已恢复回 `256`
## Phase 7: K/V block + online softmax 试验
状态:已完成并回退
尝试:
- 仅替换长序列 batched attention 路径
- 改为按 `K/V` block 流式累积的 online softmax
- 短序列和单 token 路径完全不动
目标:
- 不再一次性物化整块 `scores`
- 降低大矩阵内存压力
- 观察是否能进一步压低长段 `attention_ms`
结果:
- `/opt/qwen3-asr/temp/test_assets/qwen_rust_runtime_concurrency_2min_after_kvblock_online_softmax.json`
- `total=74.11s`
- `align=30.33s`
结论:
- 当前实现下没有优于主线最优版本
- 在这台机器上,额外的 block 循环和 online softmax 合并开销,超过了减少大 `scores` 矩阵带来的收益
- 已回退
## 当前判断
- 当前最值钱、且已验证有效的优化仍然是:
- 长序列 batched causal attention
- FFN selective F32 物化
- 继续扩大到 `QKV/O-proj` 这一步暂时不划算
- query-block 化和 batched 阈值调优目前也不划算
- `K/V block + online softmax` 在当前实现形态下也不划算
- 后续判断应优先看:
- `align_sec`
- `decoder_prefill` profiling
- 尤其是长段 `seq_len` 下的热点变化
补充:
- `asr_sec` 在多次回归中波动明显大于 `align_sec`
- 因此后续评估优化效果时,不应只盯总耗时,应优先以 `align` profiling 为准
## 暂不做的事
- 不再把 x86_64 默认路径改回 `INT8 decode`
- 不继续做没有 profiling 支撑的 `align` 结构性改写
- 不围绕 `runtime concurrency` 数量盲调
## 下一步执行顺序
1. 保留当前 batched causal attention 路径,继续观察不同长段下的稳定性
2. 保留 FFN selective F32 物化,继续观察其稳定收益
3. 如需继续优化,优先看:
- 更激进的 `attention` 算法级改动,例如按 `K/V` block 的 online softmax
- 再其次才是 `qkv_ms + gate_up_ms + down_proj_ms`
4. 如果后续继续深挖,再考虑:
- batched attention 的 block 化,降低大 score matrix 的瞬时内存
- `decoder_prefill` 内的投影层进一步收敛
5. 保持同一份 2 分钟样本持续回归,避免再次把回归误当成优化

699
docs/deployment.md 100644
View File

@ -0,0 +1,699 @@
# Qwen3-ASR 部署指南
快速部署 Qwen3-ASR 语音识别服务,支持 CPU/macOS、NVIDIA GPU、沐曦 GPU、天数 GPU 与摩尔线程 GPU 运行形态。
如果你正在继续验证本轮 CUDA 官方 vLLM 迁移,请同时参考:
- [PENDING_CUDA_VLLM_HANDOFF.md](./TODO/PENDING_CUDA_VLLM_HANDOFF.md)
依赖安装现在改成根目录默认 NVIDIA GPU,CPU、沐曦、天数与摩尔线程为单独特化环境:
| 模式 | 命令 | 说明 |
|------|------|------|
| NVIDIA GPU | `uv sync` 或 `./scripts/sync_gpu_env.sh` | Linux/NVIDIA 运行时,默认锁定 CUDA 13.0/cu130 `torch 2.11.0` / `torchaudio 2.11.0` / `torchvision 0.26.0` + `vllm 0.20.0` |
| 沐曦 GPU | `./scripts/sync_metax_env.sh` | 同步公共依赖;可选 GPU 栈安装默认从沐曦 MACA PyPI 源按 `--no-deps` 安装 |
| 天数 GPU | `./scripts/sync_iluvatar_env.sh` | 同步公共依赖;GPU 栈使用天数官方 vLLM 镜像 |
| 摩尔线程 GPU | `./scripts/sync_mthreads_env.sh` | 同步公共依赖;GPU 栈使用摩尔线程官方 MUSA vLLM 镜像 |
| CPU | `./scripts/sync_cpu_env.sh` | Linux/CPU 运行时 |
| 自动 | `./scripts/sync_accel_env.sh` | 根据 `mx-smi` / `ixsmi` / `mthreads-gmi` / `nvidia-smi` 自动选择沐曦、天数、摩尔线程、NVIDIA 或 CPU 环境 |
## 快速部署
### NVIDIA GPU 版本部署(推荐)
适用于生产环境,提供更快的推理速度:
**前置要求:**
- NVIDIA GPU(默认镜像面向 CUDA 13.0+;CUDA 12.6 / 13.0 可通过构建参数覆盖)
- 已安装 NVIDIA Container Toolkit
- 显存 12GB+(推荐 16GB+ 以支持 Qwen3-ASR 1.7B)
```bash
# 使用 docker run(带模型挂载)
docker run -d --name qwen3-asr \
--gpus all \
-p 17003:8000 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
-e ACCELERATOR=nvidia \
-e DEVICE=auto \
-e QWEN_GPU_MEMORY_UTILIZATION=0.3 \
-e QWEN_VLLM_ENFORCE_EAGER=true \
unis/qwen3-asr:gpu-latest
# 或使用 docker-compose(推荐)
docker-compose up -d
```
### 沐曦 GPU 版本部署
适用于已安装沐曦驱动与容器运行栈的机器:
```bash
docker run -d --name qwen3-asr-metax \
--privileged \
--network=host \
--pid=host \
--ipc=host \
-v /dev:/dev \
-v /opt/mxdriver:/opt/mxdriver:ro \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
-e ACCELERATOR=metax \
-e PORT=17003 \
-e METAX_VISIBLE_DEVICES=0 \
unis/qwen3-asr:metax-latest
# 或使用 docker-compose-metax.yml
docker compose -f docker-compose-metax.yml up -d
```
构建沐曦镜像时,`Dockerfile.metax` 会基于沐曦官方 vLLM 镜像融合本项目代码与通用依赖:
```bash
./scripts/package_vendor_gpu_image.sh \
--vendor metax \
--base-image <沐曦官方vLLM镜像名> \
-v n260-3.7.0.38
```
沐曦等国产 GPU 的生产推荐路径是“厂商官方 vLLM 镜像 + 本项目代码/通用依赖”。不要在项目 Dockerfile 中重新 `pip install vllm`,避免解析到 PyPI/NVIDIA CUDA 依赖。
沐曦 GPU 的完整编译、模型准备和离线交付流程见 [沐曦 GPU 国产化离线部署指南](./metax_offline_deployment.md)。
### 天数 GPU 版本部署
适用于已安装天数驱动与容器运行栈的机器。按天数官方镜像运行建议,本 compose 使用 host network、host pid/ipc、privileged、`/dev`、`/usr/src`、`/lib/modules` 等挂载,并额外挂载本项目模型与数据目录:
```bash
docker pull registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
./scripts/package_vendor_gpu_image.sh \
--vendor iluvatar \
--base-image registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5
ASR_IMAGE=unis/qwen3-asr:iluvatar-vllm0.17.0-4.4.0-v5 \
docker compose -f docker-compose-iluvatar.yml up -d
```
天数 GPU 的完整编译、模型准备和离线交付流程见 [天数 GPU 国产化离线部署指南](./iluvatar_offline_deployment.md)。
### 摩尔线程 GPU 版本部署
适用于已安装摩尔线程驱动与 MUSA 容器运行栈的机器:
```bash
docker pull registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
./scripts/package_vendor_gpu_image.sh \
--vendor mthreads \
--base-image registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519 \
-v s4000_4.3.5_d0519
ASR_IMAGE=unis/qwen3-asr:mthreads-s4000_4.3.5_d0519 \
docker compose -f docker-compose-mthreads.yml up -d
```
摩尔线程 GPU 的完整编译、模型准备和离线交付流程见 [摩尔线程 GPU 国产化离线部署指南](./mthreads_offline_deployment.md)。
默认推荐将宿主机目录统一挂载到 `/opt/dep/asr` 下,目录结构如下:
```text
/opt/dep/asr/models/
Qwen/
iic/
damo/
```
如果你希望挂载到自定义目录,可统一设置:
```bash
export MODEL_STORAGE_DIR=/data/qwen3-asr-models
export DATA_STORAGE_DIR=/data/qwen3-asr-data
mkdir -p "$MODEL_STORAGE_DIR" "$DATA_STORAGE_DIR"
docker-compose up -d
```
### 多 GPU 拓扑模式
项目现在支持统一的多 GPU 拓扑开关:
- `ASR_DEPLOY_TOPOLOGY=isolated`
- 默认模式
- 每张卡启动 1 个 backend 实例
- 容器内使用 Nginx 负载均衡到多个实例
- `ASR_DEPLOY_TOPOLOGY=sharded`
- 单个 backend 进程占用多张卡
- 由 vLLM 在进程内部做多卡分片
- `ASR_DEPLOY_TOPOLOGY=auto`
- 优先尝试 `sharded`
- 如果当前平台、可见设备或 shard 数不满足条件,则自动回退到 `isolated`
### NVIDIA 多 GPU 自动并行部署(推荐)
适用于并发量较高场景。该方案通过容器 entrypoint 自动完成:
- 根据 `ASR_VISIBLE_DEVICES` 拉起多个 ASR 实例(每张卡 1 个实例)
- 容器内自动生成 Nginx upstream 并负载均衡到各实例
- 对外仍只暴露一个服务端口(默认 `8000`)
你不需要手工维护多个 `docker-compose` 服务块或手工维护 nginx upstream。
```bash
# 4 卡示例:GPU0,1,2,3 各启动 1 个实例
ASR_VISIBLE_DEVICES=0,1,2,3 docker-compose up -d
```
常用组合:
- 单卡(保持默认):`ASR_VISIBLE_DEVICES=0`
- 双卡:`ASR_VISIBLE_DEVICES=0,1`
- 四卡:`ASR_VISIBLE_DEVICES=0,1,2,3`
强制单实例多卡分片:
```bash
ASR_DEPLOY_TOPOLOGY=sharded \
ASR_VISIBLE_DEVICES=0,1 \
docker-compose up -d
```
自动选择模式:
```bash
ASR_DEPLOY_TOPOLOGY=auto \
ASR_VISIBLE_DEVICES=0,1 \
docker-compose up -d
```
**服务访问地址:**
- API 服务: `http://localhost:17003`
- API 文档: `http://localhost:17003/docs`
### 沐曦多 GPU 自动并行部署
```bash
ASR_VISIBLE_DEVICES=0,1 docker compose -f docker-compose-metax.yml up -d
```
### 摩尔线程多 GPU 自动并行部署
```bash
ASR_VISIBLE_DEVICES=0,1 docker compose -f docker-compose-mthreads.yml up -d
```
### CPU 版本部署
适用于开发测试或无 GPU 环境:
```bash
docker run -d --name qwen3-asr \
-p 17003:8000 \
-v /opt/dep/asr/models:/app/models \
-v /opt/dep/asr/data:/app/data \
-e DEVICE=cpu \
unis/qwen3-asr:cpu-latest
```
**注意:** CPU 版本不使用 GPU/vLLM 路径。
当前 CPU 镜像已集成 QwenASR Rust backend,会自动选择 `qwen3-asr-0.6b`。
CPU 镜像默认使用可分发的 `x86-64-v2` Rust 构建目标,避免把构建机的 native CPU 指令带入通用镜像。
如果你确认构建机与部署机 CPU 指令集一致,可在自建镜像时设置 `QWENASR_RUST_TARGET_CPU=native` 换取更激进优化。
当前 Rust backend 的 x86 kernel 需要 `avx2` 与 `fma`,不满足时启动会给出明确错误。镜像默认限制
`OPENBLAS_NUM_THREADS=1` / `OMP_NUM_THREADS=1` / `GOTO_NUM_THREADS=1`,以减少多 runtime 并发时的线程争抢。
CUDA vLLM 与 CPU Rust 路径下,`word_timestamps=true` 会自动调用 forced aligner 返回字词级时间戳。
### 离线交付目录导出
如果你需要把镜像交付到不能联网的机器,推荐直接生成一个完整的离线交付目录。该脚本使用普通 `docker build` + `docker save`,不依赖 `buildx`:
```bash
# 生成 GPU 离线交付目录
./export_offline_bundle.sh --type gpu
# 或生成 CPU 离线交付目录
./export_offline_bundle.sh --type cpu
# 或生成沐曦 GPU 离线交付目录
./export_offline_bundle.sh \
--type metax \
--metax-base cr.metax-tech.com/public-ai-release/maca/vllm-metax:0.17.0-maca.ai3.5.3.307-torch2.8-py312-ubuntu22.04-amd64 \
--skip-models
./export_offline_bundle.sh --type iluvatar --iluvatar-base registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
./export_offline_bundle.sh --type mthreads --mthreads-base registry.mthreads.com/presale/devtech/vllm_musa:s4000_4.3.5_d0519
# 或一次同时生成 GPU + CPU 离线交付目录
./export_offline_bundle.sh --type all
```
生成后的目录形如:
```text
build-file/
20260520_153000-all/
qwen3-asr-cpu-20260520_153000-amd64.tar.gz
qwen3-asr-gpu-20260520_153000-amd64.tar.gz
docker-compose.yml
docker-compose-cpu.yml
docker-compose-metax.yml
.env.example
init_host_dirs.sh
README.md
DEPLOYMENT.md
```
其中会自动包含:
- 对应类型的一份或两份镜像压缩包
- 对应类型的 compose 文件;天数包只包含 `docker-compose-iluvatar.yml`
- 摩尔线程包只包含 `docker-compose-mthreads.yml`
- `.env` 模板
- 宿主机目录初始化脚本
- 离线部署说明文档
模型文件建议使用 `./scripts/download-models.sh --models-dir /opt/dep/asr/models` 单独准备;该脚本增量补齐缺失模型,不删除已有目录,也不依赖 uv。
当前运行时 / 设备默认值以主 README 为准:
- `README.md`
- `docs/README_zh.md`
设计背景与实现思路可参考:
- 当前 Qwen3 后端:`NVIDIA/沐曦 GPU -> vLLM`、`CPU/macOS -> vendored QwenASR Rust`
- 引用项目 [QwenASR](https://github.com/huanglizhuo/QwenASR)
### macOS / Apple Silicon 本地部署
适用于 M1/M2/M3/M4 机器上的本地 Qwen3-ASR 推理。当前 macOS 已统一走 vendored QwenASR Rust CPU backend。
```bash
./scripts/sync_cpu_env.sh
source .venv/bin/activate
python start.py
```
### 验证部署
```bash
# 健康检查
curl http://localhost:17003/stream/v1/asr/health
# 查看可用模型
curl http://localhost:17003/stream/v1/asr/models
# 测试语音识别(阿里云协议)
curl -X POST "http://localhost:17003/stream/v1/asr" \
-H "Content-Type: application/octet-stream" \
--data-binary @test.wav
# 测试 OpenAI 兼容接口
curl -X POST "http://localhost:17003/v1/audio/transcriptions" \
-H "Authorization: Bearer any" \
-F "file=@test.wav" \
-F "model=qwen3-asr-1.7b"
```
## 从源码构建镜像
### 使用构建脚本
项目提供了一个更薄的 `scripts/build_docker.sh` 包装层,用于统一 `docker buildx` 参数:
```bash
# 构建所有版本(CPU + GPU)
./scripts/build_docker.sh
# 仅构建 GPU 版本
./scripts/build_docker.sh -t gpu
# 构建指定版本并推送
./scripts/build_docker.sh -t all -v 1.0.1 -p
# 查看帮助
./scripts/build_docker.sh -h
```
**构建脚本参数:**
| 参数 | 说明 | 默认值 |
|------|------|--------|
| `-a, --arch` | 目标架构: `amd64`, `arm64`, `multi` | `amd64` |
| `-t, --type` | 构建类型: `cpu`, `gpu`, `all` | `all` |
| `-v, --version` | 版本标签 | `latest` |
| `-p, --push` | 构建后推送到 Docker Hub | 否 |
| `-e, --export` | 导出单架构镜像为 tar.gz | 否 |
| `-o, --output` | 导出目录 | `.` |
| `-r, --registry` | 镜像仓库 | `unis` |
| `-n, --no-cache` | 禁用 Docker 构建缓存 | 否 |
### 手动构建
```bash
# 构建 CPU 版本
docker build -t qwen3-asr:cpu-latest -f Dockerfile.cpu .
# 构建绑定当前机器指令集的 CPU 版本(仅适合同构部署)
docker build -t qwen3-asr:cpu-native -f Dockerfile.cpu \
--build-arg QWENASR_RUST_TARGET_CPU=native \
.
# 构建默认 GPU 版本(CUDA 13.0 / PyTorch cu130)
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu .
# 构建 CUDA 12.6 版本
docker build -t qwen3-asr:gpu-cu126 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda12.6-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu126 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-12-6 \
--build-arg TORCH_CUDA_ARCH_LIST="8.0;8.6;8.9" \
.
# 构建 CUDA 13.0 版本
docker build -t qwen3-asr:gpu-cu130 -f Dockerfile.gpu \
--build-arg PYTORCH_BASE_IMAGE=pytorch/pytorch:2.11.0-cuda13.0-cudnn9-runtime \
--build-arg PYTORCH_CUDA_INDEX=https://download.pytorch.org/whl/cu130 \
--build-arg CUDA_NVCC_PACKAGE=cuda-nvcc-13-0 \
--build-arg TORCH_CUDA_ARCH_LIST="12.0+PTX" \
.
```
`Dockerfile.cpu` 可覆盖的 CPU 构建参数:
| 参数 | 默认值 | 说明 |
|------|--------|------|
| `QWENASR_RUST_TARGET_CPU` | `x86-64-v2` | amd64 Rust backend 编译目标;可设为 `native` 构建绑定当前 CPU 的镜像 |
`Dockerfile.gpu` 可覆盖的 GPU 构建参数:
| 参数 | 默认值 | 用途 |
|------|--------|------|
| `PYTORCH_BASE_IMAGE` | `pytorch/pytorch:2.11.0-cuda13.0-cudnn9-runtime` | 选择 PyTorch/CUDA 基础镜像 |
| `PYTORCH_CUDA_INDEX` | `https://download.pytorch.org/whl/cu130` | 选择 PyTorch wheel CUDA 后端 |
| `CUDA_NVCC_PACKAGE` | `cuda-nvcc-13-0` | 安装匹配的 nvcc,用于 vLLM/FlashInfer JIT |
| `TORCH_CUDA_ARCH_LIST` | `12.0+PTX` | 指定 JIT 编译目标架构 |
| `VLLM_PACKAGE` | `vllm==0.20.0` | 覆盖 vLLM 包版本或来源 |
### 模型说明
服务支持以下 ASR 模型:
| 模型 | 说明 | 适用场景 |
|------|------|----------|
| Qwen3-ASR-1.7B ⭐ | 多语言 ASR(52种语言+方言,字级时间戳) | CUDA |
| Qwen3-ASR-0.6B | 轻量版多语言 ASR | CUDA / CPU Rust / macOS |
**运行时模型选择:**
系统根据机器资源自动选择合适的 Qwen3-ASR 模型:
- **显存 >= 32GB**: 自动加载 `qwen3-asr-1.7b`
- **显存 < 32GB**: 自动加载 `qwen3-asr-0.6b`
- **无 CUDA**: 自动加载基于 vendored Rust 的 `qwen3-asr-0.6b`
- **macOS / Apple Silicon**: 无论内存大小多少,默认都加载 `qwen3-asr-0.6b`
- **环境变量覆盖**: 设置 `QWEN3_ASR_MODEL=qwen3-asr-1.7b` 或 `QWEN3_ASR_MODEL=qwen3-asr-0.6b` 可硬覆盖自动选择
### 模型下载
启动时会先检测当前运行计划所需模型;如果本地缓存缺失,会自动从 ModelScope 下载。离线部署请提前准备模型缓存。
手动准备方式:
```bash
# 增量补齐离线部署所需模型,不删除已有文件
./scripts/download-models.sh --models-dir /opt/dep/asr/models
# 如果本机缺少 Python 依赖,也可以使用已构建镜像下载
ASR_IMAGE=unis/qwen3-asr:iluvatar-vllm0.17.0-4.4.0-v5 \
./scripts/download-models.sh --mode docker --models-dir /opt/dep/asr/models
```
离线部署时,推荐目录结构:
```text
/opt/dep/asr/models/
Qwen/
iic/
damo/
```
然后保持与 compose 文件一致的挂载:
```yaml
volumes:
- ${MODEL_STORAGE_DIR:-/opt/dep/asr/models}:/app/models
- ${DATA_STORAGE_DIR:-/opt/dep/asr/data}:/app/data
```
## 环境变量配置
### 基础配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `HOST` | `0.0.0.0` | 服务绑定地址 |
| `PORT` | `8000` | 服务端口 |
| `DEBUG` | `false` | 调试模式(启用后可访问 /docs) |
| `LOG_LEVEL` | `INFO` | 日志级别:DEBUG, INFO, WARNING, ERROR |
| `WORKERS` | `1` | 工作进程数(多进程会复制模型,显存成倍增加) |
| `MAX_AUDIO_SIZE` | `2048` | 最大音频文件大小(MB,支持单位如 2GB) |
| `API_KEY` | - | 服务端统一鉴权密钥 |
### 设备配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `DEVICE` | `auto` | 设备选择:`auto`, `cpu`, `cuda:0` |
| `ASR_VISIBLE_DEVICES` | `0` | 统一可见 GPU 设备配置,程序会按当前 accelerator 自动映射到底层变量 |
| `ASR_DEPLOY_TOPOLOGY` | `isolated` | 部署拓扑:`isolated`, `sharded`, `auto` |
### 内置 Nginx 与限流配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `NGINX_RATE_LIMIT_RPS` | `0` | 全局每秒请求上限,`0` 表示关闭 |
| `NGINX_RATE_LIMIT_BURST` | `0` | 全局突发请求数,`0` 时自动取 `NGINX_RATE_LIMIT_RPS` |
### ASR 模型配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `ASR_ENABLE_REALTIME_PUNC` | `true` | 是否启用实时标点模型 |
### 性能优化配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `ASR_BATCH_SIZE` | `4` | 长音频分段后的 ASR 批处理大小 |
| `INFERENCE_THREAD_POOL_SIZE` | 自动 | 推理线程池大小;默认按 CPU 核数自动设置 |
| `MAX_SEGMENT_SEC` | `60` | 音频分段最大时长(秒) |
| `QWEN_GPU_MEMORY_UTILIZATION` | `0.9` | vLLM 可保留的 GPU 显存上限,KV cache 不足时可适当调高 |
| `QWEN_VLLM_ENFORCE_EAGER` | `true` | 强制 vLLM eager 执行以提高兼容性;NVIDIA 性能测试可设为 `false` 允许 CUDA Graph 优化 |
| `WS_MAX_BUFFER_SIZE` | `160000` | WebSocket 音频缓冲区大小(样本数) |
| `QWEN_RUST_CPU_WORKERS` | `4` | CPU Rust backend worker 数;Rust ASR / forced align 默认按该数量并行 |
| `QWEN_RUST_ASR_CONCURRENCY` | `0` | Rust ASR 阶段批内并行度;`0` 表示跟随 `QWEN_RUST_CPU_WORKERS` |
| `QWEN_RUST_ALIGN_CONCURRENCY` | `0` | Rust forced align 阶段批内并行度;`0` 表示跟随 `QWEN_RUST_CPU_WORKERS` |
| `QWENASR_LIBRARY_PATH` | 自动探测 | 覆盖 vendored Rust 动态库路径 |
### 远场过滤配置
流式 ASR 远场声音过滤功能,自动过滤远场声音和环境音:
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `ASR_ENABLE_NEARFIELD_FILTER` | `true` | 启用远场声音过滤 |
| `ASR_NEARFIELD_RMS_THRESHOLD` | `0.01` | RMS 能量阈值 |
| `LOG_LEVEL=DEBUG` | - | 需要观察过滤细节时打开调试日志 |
调优建议:
- `ASR_NEARFIELD_RMS_THRESHOLD=0.01` 是当前默认值,也是推荐起点
- 嘈杂环境可以适当调高,增强背景语音过滤
- 安静环境如果出现小声说话漏识别,可以适当调低
- 需要观察过滤行为时,可临时设置 `LOG_LEVEL=DEBUG`
### 鉴权配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `API_KEY` | - | 服务端统一鉴权密钥;同时兼容 `Authorization: Bearer` 和 `X-NLS-Token` |
**使用示例:**
```bash
# 使用 Token
curl -H "X-NLS-Token: your_token" http://localhost:8000/stream/v1/asr/health
# 使用 Bearer Token(OpenAI 兼容)
curl -H "Authorization: Bearer your_token" http://localhost:8000/v1/models
```
### 日志配置
| 环境变量 | 默认值 | 说明 |
|----------|--------|------|
| `LOG_LEVEL` | `INFO` | 日志级别:`DEBUG`, `INFO`, `WARNING` |
| `LOG_FILE` | `data/logs/qwen3-asr.log` | 日志文件路径 |
| `LOG_MAX_BYTES` | `20971520` | 单个日志文件最大大小(20MB) |
| `LOG_BACKUP_COUNT` | `50` | 日志备份文件数量 |
## Docker Compose 配置
### 基础配置(GPU)
```yaml
services:
qwen3-asr:
image: unis/qwen3-asr:gpu-latest
container_name: qwen3-asr
ports:
- "17003:8000"
volumes:
- /opt/dep/asr/models:/app/models
- /opt/dep/asr/data:/app/data
environment:
- DEBUG=false
- LOG_LEVEL=INFO
- DEVICE=auto
- QWEN_GPU_MEMORY_UTILIZATION=0.3
- QWEN_VLLM_ENFORCE_EAGER=true
- ASR_BATCH_SIZE=4
- WORKERS=1
- INFERENCE_THREAD_POOL_SIZE=4
restart: unless-stopped
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: all
capabilities: [gpu]
```
### CPU 版本配置
```yaml
services:
qwen3-asr:
image: unis/qwen3-asr:cpu-latest
container_name: qwen3-asr
ports:
- "17003:8000"
volumes:
- /opt/dep/asr/models:/app/models
- /opt/dep/asr/data:/app/data
environment:
- DEBUG=false
- LOG_LEVEL=INFO
- DEVICE=cpu
- WORKERS=1
- INFERENCE_THREAD_POOL_SIZE=1
restart: unless-stopped
```
### 生产环境配置(内置 Nginx,推荐)
```yaml
services:
qwen3-asr:
image: unis/qwen3-asr:gpu-latest
container_name: qwen3-asr
ports:
- "17003:8000"
volumes:
- /opt/dep/asr/models:/app/models
- /opt/dep/asr/data:/app/data
environment:
- DEBUG=false
- LOG_LEVEL=INFO
- DEVICE=auto
- CUDA_VISIBLE_DEVICES=0,1
- QWEN_GPU_MEMORY_UTILIZATION=0.3
- QWEN_VLLM_ENFORCE_EAGER=true
- NGINX_RATE_LIMIT_RPS=20
- NGINX_RATE_LIMIT_BURST=40
- WORKERS=1
- INFERENCE_THREAD_POOL_SIZE=4
- ASR_BATCH_SIZE=4
restart: unless-stopped
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: all
capabilities: [gpu]
```
## 服务监控
### 健康检查
```bash
curl http://localhost:17003/stream/v1/asr/health
```
### 日志监控
```bash
# 实时查看日志
docker logs -f qwen3-asr
# 查看错误日志
docker logs qwen3-asr 2>&1 | grep -i error
```
### 资源监控
```bash
# 容器资源使用
docker stats qwen3-asr
# GPU 使用情况
docker exec -it qwen3-asr nvidia-smi
```
## 资源需求
### 最小配置(CPU 版本)
- CPU: 4 核
- 内存: 8GB
- 磁盘: 10GB
### 推荐配置(GPU 版本)
- CPU: 8 核
- 内存: 16GB
- GPU: NVIDIA GPU (12GB+ 显存,含说话人分离模型)
- 磁盘: 25GB
## 故障排除
### 常见问题
| 问题 | 症状 | 解决方案 |
|------|------|----------|
| GPU 内存不足 | CUDA OOM 错误 | 设置 `DEVICE=cpu` 或使用更大显存的 GPU |
| 模型加载失败 / 缓慢 | 本地模型缓存缺失 | 先运行 `./scripts/download-models.sh --models-dir /opt/dep/asr/models` 预准备模型 |
| 端口被占用 | 端口冲突错误 | 修改端口映射:`"8080:8000"` |
| 说话人分离失败 | CAM++ 模型错误 | 检查模型是否完整下载,显存是否充足 |
### 调试模式
```bash
# 启用调试模式
docker run -e DEBUG=true -e LOG_LEVEL=DEBUG ...
# 进入容器调试
docker exec -it qwen3-asr /bin/bash
```
## 更新服务
```bash
# 拉取最新镜像(GPU 版本)
docker pull unis/qwen3-asr:gpu-latest
# 拉取最新镜像(CPU 版本)
docker pull unis/qwen3-asr:cpu-latest
# 重启服务
docker-compose down && docker-compose up -d
```

View File

@ -0,0 +1,242 @@
# 天数 GPU 国产化离线部署指南
本文档用于天数/Iluvatar GPU 环境的编译、模型准备、离线包导出与目标机部署。
## 推荐基础镜像
天数默认推荐使用以下官方 vLLM 镜像作为基础镜像:
```bash
registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
```
该镜像负责提供天数 IX 运行时、PyTorch、vLLM 与相关内核。本项目的 `Dockerfile.iluvatar` 只叠加通用 Python 依赖和业务代码,不在 Dockerfile 内重新安装 vLLM,也不使用 uv 虚拟环境。
## 目录约定
目标机推荐统一使用以下宿主机目录:
```text
/opt/dep/asr/
models/
data/
logs/
temp/
tasks/
```
容器内默认挂载为:
```text
/app/models
/app/data
```
## 在线编译融合镜像
在可访问天数镜像仓库和 Python 包源的构建机上执行:
```bash
docker pull registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5
./scripts/package_vendor_gpu_image.sh \
--vendor iluvatar \
--base-image registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5
```
脚本会生成融合镜像:
```text
unis/qwen3-asr:iluvatar-vllm0.17.0-4.4.0-v5
```
并导出镜像归档到:
```text
build-file/qwen3-asr-iluvatar-vllm0.17.0-4.4.0-v5-amd64.tar.gz
```
如果只需要本机镜像,不需要导出 tar 包,可增加 `--no-export`。
## 单独准备模型
模型下载建议独立于业务服务执行。`download-models.sh` 是增量下载脚本,不会删除已有模型目录,也不依赖 uv。
在项目根目录或离线包目录执行:
```bash
./scripts/download-models.sh --models-dir /opt/dep/asr/models
```
如果本机没有 Python 依赖,但已经有融合镜像,可以用镜像内环境下载:
```bash
ASR_IMAGE=unis/qwen3-asr:iluvatar-vllm0.17.0-4.4.0-v5 \
./scripts/download-models.sh \
--mode docker \
--models-dir /opt/dep/asr/models
```
模型目录最终应至少包含:
```text
/opt/dep/asr/models/
Qwen/
Qwen3-ASR-0.6B/
Qwen3-ForcedAligner-0.6B/
damo/
iic/
```
## 导出完整离线交付包
如果要交付给不能联网的目标机,推荐直接导出天数专用离线包:
```bash
./export_offline_bundle.sh \
--type iluvatar \
--iluvatar-base registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5
```
如果模型要单独准备,不希望离线包包含模型压缩包:
```bash
./export_offline_bundle.sh \
--type iluvatar \
--iluvatar-base registry.iluvatar.com.cn:10443/customer/sz/vllm0.17.0-4.4.0-x86:v5 \
-v vllm0.17.0-4.4.0-v5 \
--skip-models
```
天数离线包只会包含天数专用启动文件:
```text
docker-compose-iluvatar.yml
.env
.env.example
init_host_dirs.sh
download-models.sh
download_models_standalone.py
README.md
DEPLOYMENT.md
BUNDLE_INFO.txt
qwen3-asr-iluvatar-*-amd64.tar.gz
```
不会再要求使用通用 `docker-compose.yml`。
## 目标机离线部署
将整个离线包复制到目标机后,进入离线包目录:
```bash
chmod +x init_host_dirs.sh download-models.sh
./init_host_dirs.sh
```
导入镜像:
```bash
gunzip -c qwen3-asr-iluvatar-vllm0.17.0-4.4.0-v5-amd64.tar.gz | docker load
```
确认 `.env` 中的镜像名与导入镜像一致:
```env
ASR_IMAGE=unis/qwen3-asr:iluvatar-vllm0.17.0-4.4.0-v5
```
按需设置显卡与 vLLM 显存比例:
```env
ILUVATAR_VISIBLE_DEVICES=0
IX_VISIBLE_DEVICES=0
CUDA_VISIBLE_DEVICES=0
QWEN_GPU_MEMORY_UTILIZATION=0.25
QWEN_VLLM_ENFORCE_EAGER=true
```
启动服务:
```bash
docker compose -f docker-compose-iluvatar.yml up -d
```
查看状态与日志:
```bash
docker compose -f docker-compose-iluvatar.yml ps
docker compose -f docker-compose-iluvatar.yml logs -f
```
服务默认监听 host 网络端口:
```text
http://<目标机IP>:17003
```
## 已导出的旧离线包处理
如果旧离线包里的镜像已经能识别 `Qwen3ASRForConditionalGeneration`,但启动时报 KV cache 不足,例如:
```text
Try increasing gpu_memory_utilization or decreasing max_model_len
```
不需要重新打镜像。只需要在旧离线包的 `docker-compose-iluvatar.yml` 的 `QWEN3_ASR_MODEL` 附近补充:
```yaml
QWEN_GPU_MEMORY_UTILIZATION: ${QWEN_GPU_MEMORY_UTILIZATION:-0.25}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
```
然后重新创建容器:
```bash
docker compose -f docker-compose-iluvatar.yml down
docker compose -f docker-compose-iluvatar.yml up -d
```
只执行 `restart` 不会重新注入环境变量。
## 常见问题
### 为什么不用 uv?
天数 Docker 镜像内推荐直接使用系统 Python 环境。`Dockerfile.iluvatar` 使用:
```bash
python3 -m pip install --no-cache-dir -r environments/iluvatar/requirements.txt
```
不创建 uv 虚拟环境,也不在镜像内运行 `uv pip install`。
### 为什么不重新安装 vLLM?
国产 GPU 的 vLLM、PyTorch、内核和运行时通常需要严格匹配厂商镜像。项目层重新 `pip install vllm` 容易解析到 PyPI/NVIDIA CUDA 依赖,破坏天数官方镜像里的匹配关系。
### `QWEN_GPU_MEMORY_UTILIZATION` 为什么没生效?
变量必须进入容器才会生效。天数 compose 中需要有:
```yaml
QWEN_GPU_MEMORY_UTILIZATION: ${QWEN_GPU_MEMORY_UTILIZATION:-}
QWEN_VLLM_ENFORCE_EAGER: ${QWEN_VLLM_ENFORCE_EAGER:-true}
```
然后在 `.env` 设置:
```env
QWEN_GPU_MEMORY_UTILIZATION=0.25
QWEN_VLLM_ENFORCE_EAGER=true
```
修改 `.env` 后必须 `down` 再 `up -d`。
### 日志中的本地模型 repo id warning 是否致命?
vLLM 可能会先尝试按远端 repo 方式读取 safetensors,遇到本地路径时打印 warning。如果后续出现 `Loading safetensors checkpoint shards` 并继续加载权重,通常不是致命错误。
真正需要处理的是最后的异常,例如模型架构不识别、KV cache 不足、模型文件缺失等。

Some files were not shown because too many files have changed in this diff Show More