合并主干
parent
5c58759099
commit
71f8d572c0
41
.env.example
41
.env.example
|
|
@ -1,41 +1,52 @@
|
|||
# Local model assets; the backend never downloads models at startup.
|
||||
# 本地模型目录;服务启动时不会自动下载模型权重。
|
||||
MODEL_DIR=models
|
||||
# Set this if CAM++ is not at a model_manifest.json path.
|
||||
# CAM++ 模型不在 model_manifest.json 默认目录时,可用此项指定其他本地路径。
|
||||
# CAM_MODEL_PATH=iic/speech_campplus_sv_zh-cn_16k-common
|
||||
|
||||
# FunASR realtime engine
|
||||
# FunASR 流式 ASR 与 VAD 模型。
|
||||
FUNASR_ASR_MODEL=paraformer-zh-streaming
|
||||
FUNASR_VAD_MODEL=fsmn-vad
|
||||
FUNASR_DEVICE=cuda:0
|
||||
FUNASR_VAD_DEVICE=cpu
|
||||
# FunASR 原生 WebSocket 服务使用的 CPU 工作线程数。
|
||||
FUNASR_NCPU=4
|
||||
|
||||
# VAD 结束语音段前等待的静音时长:句子模式为 800 毫秒,段落模式默认保留 5 秒。
|
||||
FUNASR_VAD_MAX_END_SILENCE_MS=800
|
||||
FUNASR_VAD_PARAGRAPH_MAX_END_SILENCE_MS=5000
|
||||
# 提高此阈值会过滤更多环境噪声,也可能漏掉较轻的语音。
|
||||
FUNASR_VAD_SPEECH_NOISE_THRES=0.6
|
||||
|
||||
# FunASR 原生 WebSocket 流式识别参数。
|
||||
FUNASR_CHUNK_SIZE=0,10,5
|
||||
FUNASR_CHUNK_INTERVAL=10
|
||||
FUNASR_ENCODER_LOOK_BACK=4
|
||||
FUNASR_DECODER_LOOK_BACK=1
|
||||
FUNASR_VAD_CHUNK_MS=200
|
||||
FUNASR_MAX_SEGMENT_SEC=30
|
||||
# CAM++ 段内说话人追踪沿用 FunASR 的重叠短窗方案:1.5 秒窗、0.75 秒步长。
|
||||
FUNASR_FINALIZE_TIMEOUT_SECONDS=300
|
||||
FUNASR_NATIVE_WS_HOST=127.0.0.1
|
||||
FUNASR_NATIVE_WS_PORT=10095
|
||||
|
||||
# CAM++ 说话人聚类和会话身份参数。
|
||||
FUNASR_SPEAKER_WINDOW_MS=1500
|
||||
FUNASR_SPEAKER_HOP_MS=750
|
||||
# CAM++ 每批最多计算多少个窗口,避免很长 turn 一次性占用过多显存。
|
||||
FUNASR_SPEAKER_BATCH_SIZE=16
|
||||
# FunASR HybridSpeakerTracker history and stable identity matching defaults.
|
||||
FUNASR_SPEAKER_HISTORY_CHUNKS=128
|
||||
FUNASR_SPEAKER_MATCH_THRESHOLD=0.6
|
||||
FUNASR_SPEAKER_MERGE_THRESHOLD=0.78
|
||||
# FunASR's default speaker identity limit; increase this value for larger meetings.
|
||||
FUNASR_MAX_SPEAKERS=15
|
||||
# 小于该时长的相邻说话人段会合并,避免短窗噪声把一句话切得过碎。
|
||||
FUNASR_SPEAKER_MIN_SEGMENT_MS=3000
|
||||
|
||||
# CAM++ is required and started by the backend launcher.
|
||||
# CAM++ 为必需模型。若需将 GPU 留给 ASR,可将 AUXILIARY_DEVICE 设为 cpu。
|
||||
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010
|
||||
AUXILIARY_DEVICE=cuda:0
|
||||
AUXILIARY_PRELOAD_KINDS=speaker_verification
|
||||
# 标点模型为可选模型,默认在 CPU 上运行。
|
||||
FUNASR_PUNC_DEVICE=cpu
|
||||
|
||||
# Standalone frontend and browser-facing backend origin.
|
||||
# 前后端分别启动;前端会将 /ws 和 /api/stop 请求转发给后端。
|
||||
FRONTEND_HOST=127.0.0.1
|
||||
FRONTEND_PORT=8080
|
||||
FRONTEND_ORIGIN=http://127.0.0.1:8080
|
||||
BACKEND_PUBLIC_URL=http://127.0.0.1:8082
|
||||
|
||||
BACKEND_INTERNAL_URL=http://127.0.0.1:8082
|
||||
WEB_HOST=0.0.0.0
|
||||
WEB_PORT=8082
|
||||
WEB_DISPLAY_HOST=127.0.0.1
|
||||
|
|
|
|||
|
|
@ -1,42 +0,0 @@
|
|||
# Local model assets; model loading never downloads weights at startup.
|
||||
MODEL_DIR=models
|
||||
# Set this if CAM++ is not under the path declared in model_manifest.json.
|
||||
# CAM_MODEL_PATH=iic/speech_campplus_sv_zh-cn_16k-common
|
||||
|
||||
# Local FunASR streaming ASR and VAD models.
|
||||
FUNASR_ASR_MODEL=paraformer-zh-streaming
|
||||
FUNASR_VAD_MODEL=fsmn-vad
|
||||
FUNASR_DEVICE=cuda:0
|
||||
FUNASR_VAD_DEVICE=cpu
|
||||
# Silence duration in milliseconds before FSMN VAD finalizes a speech segment.
|
||||
# Strategy 0 (semantic sentence) defaults to 800 ms; increase to preserve pauses.
|
||||
FUNASR_VAD_MAX_END_SILENCE_MS=800
|
||||
# Strategy 1 (paragraph) keeps the longer 5 s pause. Adjust independently if needed.
|
||||
FUNASR_VAD_PARAGRAPH_MAX_END_SILENCE_MS=5000
|
||||
# Higher margins reject more weak background noise but can suppress quiet speech.
|
||||
FUNASR_VAD_SPEECH_NOISE_THRES=0.6
|
||||
|
||||
# Native FunASR WSS chunk settings. The middle chunk is sent as 10 x 60 ms.
|
||||
FUNASR_CHUNK_SIZE=0,10,5
|
||||
FUNASR_CHUNK_INTERVAL=10
|
||||
FUNASR_ENCODER_LOOK_BACK=4
|
||||
FUNASR_DECODER_LOOK_BACK=1
|
||||
FUNASR_FINALIZE_TIMEOUT_SECONDS=300
|
||||
FUNASR_NATIVE_WS_HOST=127.0.0.1
|
||||
FUNASR_NATIVE_WS_PORT=10095
|
||||
|
||||
# CAM++ is required and started by backend/run_backend.py.
|
||||
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010
|
||||
AUXILIARY_DEVICE=cpu
|
||||
AUXILIARY_PRELOAD_KINDS=speaker_verification
|
||||
# Optional punctuation runs on CPU by default so the speaker model can retain GPU memory.
|
||||
FUNASR_PUNC_DEVICE=cpu
|
||||
|
||||
# Frontend and backend run independently. The frontend proxies /ws and /api/stop.
|
||||
FRONTEND_HOST=127.0.0.1
|
||||
FRONTEND_PORT=8080
|
||||
BACKEND_INTERNAL_URL=http://127.0.0.1:8082
|
||||
BACKEND_PUBLIC_URL=http://127.0.0.1:8082
|
||||
WEB_HOST=0.0.0.0
|
||||
WEB_PORT=8082
|
||||
WEB_DISPLAY_HOST=127.0.0.1
|
||||
|
|
@ -1,50 +0,0 @@
|
|||
# 选择已经下载到 demo/models 目录中的 ASR 模型;必须与实际模型目录保持一致。
|
||||
# 默认使用轻量的 0.6B 模型。如下载的是 1.7B,请改为 1.7b,不能混用。
|
||||
QWEN3_ASR_MODEL=0.6b
|
||||
# 本地模型根目录。下载脚本、VLLM 启动脚本和辅助服务都会从这里查找模型。
|
||||
MODEL_DIR=D:/github-project/ASR/Qwen-Asr/demo/models
|
||||
|
||||
# VLLM 宿主机服务绑定地址;0.0.0.0 表示允许服务器网卡接收外部请求。
|
||||
VLLM_HOST=0.0.0.0
|
||||
# VLLM 服务端口;启动器、健康检查和前端连接地址必须使用同一个端口。
|
||||
VLLM_PORT=9950
|
||||
# VLLM 可执行文件名称;新版环境通常为 vllm,启动器会自动执行 vllm serve。
|
||||
VLLM_EXECUTABLE=vllm
|
||||
# 启动成功提示中显示的地址,只影响日志和使用说明,不改变实际监听地址。
|
||||
VLLM_DISPLAY_HOST=127.0.0.1
|
||||
# 启动轮询使用的地址;如果 VLLM 部署在本机,通常保持 127.0.0.1 即可。
|
||||
VLLM_PROBE_HOST=127.0.0.1
|
||||
# VLLM 启动阶段最多检查多少次健康状态,超过次数仍未就绪则启动失败。
|
||||
VLLM_STARTUP_CHECK_LOOPS=60
|
||||
# 两次健康检查之间的等待秒数;模型加载较慢时可以适当增大检查次数。
|
||||
VLLM_STARTUP_CHECK_INTERVAL_SECONDS=2
|
||||
# VLLM 使用的显存比例。显存还要留给 VAD、聚类和声纹辅助模型时,建议预留余量。
|
||||
VLLM_GPU_MEMORY_UTILIZATION=0.3
|
||||
# VLLM 的最大上下文长度;数值越大通常占用越多显存,请结合显卡容量调整。
|
||||
VLLM_MAX_MODEL_LEN=16384
|
||||
# VLLM 同时处理的最大序列数;实时单路验证可保持默认值,多路并发时再调大。
|
||||
VLLM_MAX_NUM_SEQS=16
|
||||
# 张量并行 GPU 数量;单卡部署为 1,多卡部署时填写参与并行的 GPU 数量。
|
||||
VLLM_TENSOR_PARALLEL_SIZE=1
|
||||
# 是否启用 eager 模式。true 通常更容易启动和排查,false 可能获得更高性能。
|
||||
VLLM_ENFORCE_EAGER=true
|
||||
|
||||
# GB10 需要使用 CUDA 13 工具链中的 ptxas;serve.py 会自动加载该配置并传给 VLLM。
|
||||
TRITON_PTXAS_PATH=/usr/local/cuda/bin/ptxas
|
||||
|
||||
# 可选:为外部客户端设置稳定的公开模型名称。留空时默认使用完整 ModelScope ID。
|
||||
# VLLM_SERVED_MODEL_NAME=Qwen/Qwen3-ASR-0.6B
|
||||
|
||||
# 辅助模型服务地址;它独立加载 VAD、CAM++ 聚类和声纹模型,WebSocket
|
||||
# 服务只通过 HTTP 调用,不会把这些 GPU 模型重复加载到 WebSocket 进程。
|
||||
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010
|
||||
# WebSocket 调用已独立部署的 vLLM;跨服务器时改成模型服务器的实际地址。
|
||||
MODEL_SERVICE_URL=http://127.0.0.1:9950/v1
|
||||
# 辅助服务监听的 GPU;单卡服务器保持 cuda:0,多卡时可改成指定卡号。
|
||||
AUXILIARY_DEVICE=cuda:0
|
||||
# 仅保留兼容旧配置;VAD/CAM++ 核心模型缺失时始终拒绝启动,不会静默降级。
|
||||
# 可选 diarization/aligner 缺失不会阻断实时服务。
|
||||
AUXILIARY_ALLOW_MISSING=false
|
||||
# 启动时强制加载 VAD + CAM++ speaker_verification;完整 diarization 首次调用时按需加载。
|
||||
# 如确实需要启动时额外预加载,可追加:vad,speaker_verification,diarization
|
||||
AUXILIARY_PRELOAD_KINDS=vad,speaker_verification
|
||||
|
|
@ -1,42 +1,42 @@
|
|||
# FunASR realtime browser demo
|
||||
# FunASR 实时识别说明
|
||||
|
||||
The frontend and backend start separately. The backend starts these supervised processes:
|
||||
前端和后端分别启动。后端会托管以下服务:
|
||||
|
||||
- FunASR native online WebSocket server, using the local streaming ASR and FSMN VAD models. The browser adapter sends FunASR's `mode=online`, chunk/look-back settings, fixed 60 ms PCM frames, and `is_speaking=false` end-of-input flush.
|
||||
- CAM++ auxiliary service, which assigns stable speaker labels to finalized utterances.
|
||||
- A small browser protocol adapter that translates the unchanged Tencent demo message format to FunASR's native WS format. It does not run a second ASR/VAD segmentation pipeline.
|
||||
- FunASR 原生在线 WebSocket 服务,使用本地流式 ASR 和 FSMN-VAD 模型。浏览器桥接层发送 `mode=online`、分块和回看参数,并在输入结束时发送 `is_speaking=false`,让引擎刷新最后的识别结果。
|
||||
- CAM++ 辅助服务,为已完成的语音段分配稳定的说话人标签。
|
||||
- 浏览器协议桥接服务,将腾讯演示页面原有的消息格式转换为 FunASR WebSocket 协议。VAD 切分由 FunASR 处理,桥接层不会再运行另一套 ASR/VAD 切分流程。
|
||||
|
||||
The frontend serves the exact files from the local `tencent-demo/static` directory and proxies the page's same-origin `/ws` and `/api/stop` requests to the backend.
|
||||
前端继续提供 `frontend/static/` 中的腾讯演示文件,并将同源的 `/ws` 和 `/api/stop` 请求转发给后端。
|
||||
|
||||
## Model directories
|
||||
## 模型目录
|
||||
|
||||
Put assets under `models/` or set `MODEL_DIR` in `.env`. The default names resolve to these local directories:
|
||||
将模型放入 `models/`,或在 `.env` 中设置 `MODEL_DIR`。默认模型目录如下:
|
||||
|
||||
- `models/iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online`
|
||||
- `models/damo/speech_fsmn_vad_zh-cn-16k-common-pytorch`
|
||||
- `models/iic/speech_campplus_sv_zh-cn_16k-common` (or the configured `damo/` variant)
|
||||
- `models/iic/speech_campplus_sv_zh-cn_16k-common`,也可使用清单中配置的 `damo/` 模型目录
|
||||
|
||||
If directories have different names, set `FUNASR_ASR_MODEL` and `FUNASR_VAD_MODEL` to their local paths. CAM++ follows `model_manifest.json`; set `CAM_MODEL_PATH` for another location. Startup checks all three assets before exposing the browser bridge.
|
||||
若模型目录名称不同,可在 `.env` 中将 `FUNASR_ASR_MODEL` 和 `FUNASR_VAD_MODEL` 设为本地路径。CAM++ 默认按 `model_manifest.json` 查找;需要自定义目录时设置 `CAM_MODEL_PATH`。后端会在对外提供 WebSocket 服务前检查 ASR、VAD 和 CAM++ 三个必需模型。
|
||||
|
||||
## Install and start
|
||||
## 安装与启动
|
||||
|
||||
The root requirements.txt includes FunASR and auxiliary-service dependencies plus the Qwen vLLM stack pinned for a GB10 CUDA 13 host. Adjust those CUDA-specific pins for other hosts before installing:
|
||||
`requirements.txt` 包含 FunASR、辅助服务和前端依赖。其中 PyTorch 版本面向使用 CUDA 13 的 GB10 环境;其它设备请按其 CUDA 版本调整 PyTorch 软件源和版本。
|
||||
|
||||
~~~powershell
|
||||
python -m pip install -r requirements.txt
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.example .env }
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.funasr.example .env }
|
||||
~~~
|
||||
|
||||
Start the backend and frontend in separate terminals:
|
||||
在两个终端中分别启动后端和前端:
|
||||
|
||||
~~~powershell
|
||||
python backend\run_backend.py
|
||||
python frontend\run_frontend.py
|
||||
~~~
|
||||
|
||||
Defaults are frontend port 8080, browser backend port 8082, CAM++ HTTP port 8010, and native FunASR WS port 10095 bound to loopback. Change `FRONTEND_PORT`, `WEB_PORT`, `AUXILIARY_SERVICE_URL`, `FUNASR_NATIVE_WS_HOST`, and `FUNASR_NATIVE_WS_PORT` together when needed. `BACKEND_INTERNAL_URL` is the backend origin reachable from the frontend process; it defaults to `http://127.0.0.1:${WEB_PORT}`.
|
||||
默认前端端口为 `8080`,浏览器后端端口为 `8082`,CAM++ HTTP 端口为 `8010`,仅本机可访问的 FunASR WebSocket 端口为 `10095`。需要调整时,在 `.env` 中同步修改 `FRONTEND_PORT`、`WEB_PORT`、`AUXILIARY_SERVICE_URL`、`FUNASR_NATIVE_WS_HOST` 和 `FUNASR_NATIVE_WS_PORT`。前端进程通过 `BACKEND_INTERNAL_URL` 访问后端;未设置时使用 `http://127.0.0.1:${WEB_PORT}`。
|
||||
|
||||
`FUNASR_DEVICE` and `FUNASR_VAD_DEVICE` control ASR and VAD placement independently. Set `AUXILIARY_DEVICE=cpu` when CAM++ should not share the ASR GPU. `FUNASR_CHUNK_SIZE` and `FUNASR_CHUNK_INTERVAL` control FunASR's native chunk protocol; the default `[0,10,5]` and interval 10 send the current chunk in 600 ms groups.
|
||||
`FUNASR_DEVICE` 和 `FUNASR_VAD_DEVICE` 分别控制 ASR 与 VAD 的运行设备。若 CAM++ 不与 ASR 共用 GPU,可将 `AUXILIARY_DEVICE` 设为 `cpu`。`FUNASR_CHUNK_SIZE` 和 `FUNASR_CHUNK_INTERVAL` 控制 FunASR 原生分块协议;默认值 `[0,10,5]` 和间隔 `10` 会按 600 毫秒一组发送当前音频块。
|
||||
|
||||
Real-time file input currently accepts raw PCM16 or 16 kHz mono PCM WAV, matching the backend's available decoder. Speaker labels are computed for every finalized utterance; very short or silent segments can still be marked as unknown by CAM++.
|
||||
实时文件输入支持原始 PCM16 或 16 kHz 单声道 PCM WAV。每个完成的语音段都会计算说话人标签;语音过短或接近静音时,CAM++ 仍可能将说话人标为未知。
|
||||
|
|
|
|||
28
README.md
28
README.md
|
|
@ -1,32 +1,28 @@
|
|||
# FunASR realtime ASR demo
|
||||
# FunASR 实时语音识别演示
|
||||
|
||||
The project keeps the frontend and backend in separate directories. The frontend serves the original Tencent demo assets from frontend/static/ unchanged.
|
||||
本项目将 FunASR 流式 WebSocket 后端与腾讯演示前端分开运行。前端沿用原有页面文件,界面和消息协议保持不变。
|
||||
|
||||
## Project layout
|
||||
## 安装与准备模型
|
||||
|
||||
- backend/: FunASR WebSocket adapter, native engine launcher, CAM++ service, and backend code.
|
||||
- frontend/: Tencent demo files and the static HTTP/WebSocket proxy.
|
||||
- scripts/: shared model download tools.
|
||||
- requirements.txt: combined dependencies for both ASR engines and the frontend.
|
||||
- model_manifest.json: shared model registry.
|
||||
|
||||
## Install and download models
|
||||
|
||||
The combined requirements file includes the GB10 CUDA 13 PyTorch and vLLM stack used by the Qwen3-ASR branch. Adjust the CUDA index and torch-family pins before installing on a different host.
|
||||
`requirements.txt` 中的 PyTorch 版本面向使用 CUDA 13 的 GB10 环境。部署到其它设备时,请按对应 CUDA 版本调整 PyTorch 软件源和版本。
|
||||
|
||||
~~~powershell
|
||||
python -m pip install -r requirements.txt
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.funasr.example .env }
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.example .env }
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
|
||||
## Start
|
||||
`--funasr-runtime` 会下载流式 Paraformer、FSMN-VAD、CAM++ 和标点模型。若要下载 `model_manifest.json` 中列出的所有可选模型,请运行 `python scripts/download_models.py`。
|
||||
|
||||
Run the backend and frontend in separate terminals from the project root:
|
||||
## 启动
|
||||
|
||||
在项目根目录打开两个终端,分别运行:
|
||||
|
||||
~~~powershell
|
||||
python backend/run_backend.py
|
||||
python frontend/run_frontend.py
|
||||
~~~
|
||||
|
||||
The backend supervises CAM++, FunASR native WSS, and the browser protocol adapter. The frontend serves the unchanged Tencent page and proxies its same-origin /ws and /api/stop requests to the backend. Configure ports and model locations in .env.
|
||||
后端负责启动 FunASR 原生 WebSocket 服务、CAM++ 服务和浏览器协议桥接服务。前端提供腾讯演示页面,并将 `/ws` 和 `/api/stop` 请求转发给后端。模型路径、设备、VAD 静音时长、说话人聚类参数和端口均可在 `.env` 中配置。
|
||||
|
||||
模型目录和运行细节见 [FunASR 使用说明](FUNASR_README.md)。
|
||||
|
|
|
|||
|
|
@ -1,151 +0,0 @@
|
|||
# Qwen3-ASR VLLM 独立部署项目
|
||||
|
||||
本目录是后续实时 ASR 功能验证使用的独立模型服务项目。
|
||||
|
||||
它不导入、不启动、也不调用仓库根目录下原项目的 `app/` 代码。模型下载、VLLM 启动、配置和服务验证都在本目录内完成。后续验证 demo 只需要调用这里提供的 VLLM OpenAI 兼容接口。
|
||||
|
||||
## 当前下载范围
|
||||
|
||||
默认下载一个 ASR 模型和独立辅助模型运行服务所需的全部模型资产:
|
||||
|
||||
```text
|
||||
Qwen/Qwen3-ASR-0.6B
|
||||
```
|
||||
|
||||
ASR 如需使用大模型,可显式选择 1.7B:
|
||||
|
||||
```text
|
||||
Qwen/Qwen3-ASR-1.7B
|
||||
```
|
||||
|
||||
辅助资产包括 VAD、CAM++ 分离、配置声纹、实时声纹、CAM++ Transformer 和 Qwen3 ForcedAligner。它们不会由 `qwen-asr-serve` 启动,而是由独立 Python 辅助模型服务加载。
|
||||
|
||||
不会下载另一个未选择的 ASR 模型;辅助模型资产会随默认部署包下载,供独立 Python 运行服务预加载。
|
||||
|
||||
## 1. 下载模型
|
||||
|
||||
模型直接下载到宿主机的 `demo/models`。ASR 由 VLLM 启动,辅助模型由独立 Python 运行服务启动。
|
||||
|
||||
```powershell
|
||||
cd D:\github-project\ASR\Qwen-Asr\demo
|
||||
python -m venv .venv
|
||||
.\.venv\Scripts\Activate.ps1
|
||||
pip install -r requirements.txt
|
||||
python scripts\download_models.py
|
||||
```
|
||||
|
||||
上述命令会下载默认 `0.6B` ASR 以及全部辅助模型,并在下载完成后把 CAM++ 配置中的依赖模型 ID 改为 `demo/models` 下的本地路径,保证辅助服务可以离线启动。只下载 ASR 时使用:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --skip-auxiliary
|
||||
```
|
||||
|
||||
只下载辅助模型时使用:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --auxiliary-only
|
||||
```
|
||||
|
||||
选择 1.7B:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --model 1.7b
|
||||
```
|
||||
|
||||
检查模型是否完整但不下载:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --check-only
|
||||
```
|
||||
|
||||
ModelScope 下载也可以通过环境变量调整缓存目录:
|
||||
|
||||
```powershell
|
||||
$env:MODELSCOPE_CACHE = 'D:\modelscope-cache'
|
||||
python scripts\download_models.py
|
||||
```
|
||||
|
||||
## 2. 宿主机启动 VLLM 服务
|
||||
|
||||
需要宿主机具备与 VLLM 兼容的 Python、CUDA 和 NVIDIA 驱动环境。安装部署依赖:
|
||||
|
||||
```bash
|
||||
python -m pip install -r requirements.txt
|
||||
```
|
||||
|
||||
先复制并按服务器实际路径修改 `.env`,启动器会自动读取该文件:
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
```
|
||||
|
||||
默认启动 `Qwen/Qwen3-ASR-0.6B`,监听地址为 `0.0.0.0:9950`:
|
||||
|
||||
```bash
|
||||
python backend/serve_qwen_legacy.py
|
||||
```
|
||||
|
||||
模型和启动检查循环可以通过命令行或环境变量传入;端口统一在 `backend/serve_qwen_legacy.py` 的 `SERVER_PORT` 变量中维护:
|
||||
|
||||
```bash
|
||||
QWEN3_ASR_MODEL=0.6b VLLM_STARTUP_CHECK_LOOPS=120 \
|
||||
VLLM_STARTUP_CHECK_INTERVAL_SECONDS=2 python backend/serve_qwen_legacy.py
|
||||
```
|
||||
|
||||
如果使用 `1.7B`,下载和启动必须指定同一个模型:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --model 1.7b
|
||||
python -m backend.serve_qwen_legacy --model 1.7b
|
||||
```
|
||||
|
||||
启动器会按 `VLLM_STARTUP_CHECK_LOOPS` 次数轮询 `/health`,每次间隔由 `VLLM_STARTUP_CHECK_INTERVAL_SECONDS` 指定。服务端口、健康检查端口和就绪提示统一使用 `backend/serve_qwen_legacy.py` 中的 `SERVER_PORT`。
|
||||
|
||||
服务启动后可检查:
|
||||
|
||||
```bash
|
||||
curl "http://${VLLM_DISPLAY_HOST:-127.0.0.1}:9950/health"
|
||||
curl "http://${VLLM_DISPLAY_HOST:-127.0.0.1}:9950/v1/models"
|
||||
```
|
||||
|
||||
## 3. 调用转写接口
|
||||
|
||||
VLLM 服务提供 OpenAI 兼容的音频转写接口:
|
||||
|
||||
```bash
|
||||
curl "http://${VLLM_DISPLAY_HOST:-127.0.0.1}:9950/v1/audio/transcriptions" \
|
||||
-H "Authorization: Bearer EMPTY" \
|
||||
-F "file=@./audio/sample.wav" \
|
||||
-F "model=Qwen/Qwen3-ASR-0.6B"
|
||||
```
|
||||
|
||||
## 4. 启动辅助模型服务
|
||||
|
||||
另开一个终端,在同一台服务器启动 VAD、CAM++ 和声纹模型运行服务:
|
||||
|
||||
```powershell
|
||||
cd D:\github-project\ASR\Qwen-Asr\demo
|
||||
pip install -r requirements.txt
|
||||
python -m backend.auxiliary_server
|
||||
```
|
||||
|
||||
辅助服务默认监听 `0.0.0.0:8010`。实时链路启动时严格加载 VAD 和 CAM++ `speaker_verification` 声纹模型,用于每个 turn 的特征提取与在线聚类;完整 CAM++ 分离、Transformer 和 ForcedAligner 不阻断核心服务,完整分离模型会在调用 `/v1/diarization` 时按需加载。检查状态:
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:8010/health
|
||||
```
|
||||
|
||||
WebSocket demo 默认连接 `9950` 的 ASR VLLM,辅助服务使用 `8010`。`/health` 的 `ready` 要求 `vad_ready` 与 `speaker_embedding_ready` 同时为 true;完整 diarization 资产缺失不会影响实时 `/v1/speaker/resolve`。
|
||||
|
||||
实时链路中的职责是:WebSocket 用 RMS 帧门控快速检测停顿;辅助服务用 FunASR VAD 提供 `/v1/vad`,并加载 CAM++ `speaker_verification` 提取 turn embedding,再由服务端在线聚类。`speech_campplus_speaker-diarization_common` 是完整音频分离接口的额外 pipeline,不是实时 turn 聚类的唯一入口。
|
||||
|
||||
也可以使用多模态 Chat Completions 接口,后续实时验证项目将以此服务边界为准。
|
||||
|
||||
## 项目边界
|
||||
|
||||
- `scripts/download_models.py`:下载选定 ASR 和全部辅助模型资产。
|
||||
- `backend/serve_qwen_legacy.py`:读取 `.env`,解析模型选择、宿主机参数并启动新版 `vllm serve`。
|
||||
- `requirements.txt`:安装宿主机部署所需的官方 Qwen3-ASR VLLM 依赖。
|
||||
- `tests/`:只验证本项目自己的模型清单和选择逻辑,不依赖原项目。
|
||||
|
||||
模型服务就绪后,新的实时 ASR demo 放在同级 `demo` 项目中继续开发,但不得通过 Python import 或 HTTP/WebSocket 调用原项目服务。
|
||||
|
|
@ -1 +1 @@
|
|||
"""Backend services and the realtime ASR protocol adapter."""
|
||||
"""后端服务与实时 ASR 协议适配器。"""
|
||||
|
|
|
|||
|
|
@ -27,18 +27,18 @@ from backend.model_manifest import auxiliary_models, load_manifest, model_direct
|
|||
# 将辅助服务端口固定在代码变量中,服务器启动时只需执行脚本,便于部署和排查。
|
||||
AUXILIARY_HOST = "0.0.0.0"
|
||||
AUXILIARY_PORT = 8010
|
||||
# 独立启动辅助服务也必须读取部署配置,不能只在启动 vLLM 时才加载 .env。
|
||||
# 独立启动辅助服务时也加载部署配置,确保模型目录和设备与启动器一致。
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
MODELS_DIR = Path(os.getenv("MODEL_DIR", str(PROJECT_ROOT / "models"))).resolve()
|
||||
AUXILIARY_DEVICE = os.getenv("AUXILIARY_DEVICE", "cuda:0")
|
||||
# Keep the real-time tracker defaults aligned with FunASR HybridSpeakerTracker.
|
||||
# 实时追踪器的默认参数与 FunASR HybridSpeakerTracker 保持一致。
|
||||
ONLINE_SPEAKER_MATCH_THRESHOLD = float(os.getenv("FUNASR_SPEAKER_MATCH_THRESHOLD", "0.6"))
|
||||
ONLINE_CLUSTER_MERGE_THRESHOLD = float(os.getenv("FUNASR_SPEAKER_MERGE_THRESHOLD", "0.78"))
|
||||
ONLINE_MAX_SPEAKERS = max(1, int(os.getenv("FUNASR_MAX_SPEAKERS", "15")))
|
||||
ONLINE_SPEAKER_HISTORY_CHUNKS = max(1, int(os.getenv("FUNASR_SPEAKER_HISTORY_CHUNKS", "128")))
|
||||
MIN_ONLINE_SPEAKER_AUDIO_MS = 800
|
||||
# Realtime VAD is loaded by the WebSocket process. This service loads CAM++.
|
||||
# Its standalone VAD HTTP endpoint remains available on demand.
|
||||
# 实时 VAD 由 WebSocket 进程加载;本服务负责加载 CAM++。
|
||||
# 独立的 VAD HTTP 接口仍可按需使用。
|
||||
DEFAULT_PRELOAD_KINDS = {"speaker_verification"}
|
||||
|
||||
|
||||
|
|
@ -97,9 +97,9 @@ class AuxiliaryRuntime:
|
|||
self.models: dict[str, Any] = {}
|
||||
self.status: dict[str, dict[str, Any]] = {}
|
||||
self.inference_lock = asyncio.Lock()
|
||||
# 每个 WebSocket session 独立维护聚类中心,避免不同浏览器会话互相污染。
|
||||
# 每个 WebSocket 会话独立维护聚类中心,避免不同浏览器会话互相影响。
|
||||
self.speaker_clusters: dict[str, list[dict[str, Any]]] = {}
|
||||
# Keep a bounded rolling window history per WebSocket session, like FunASR's tracker.
|
||||
# 像 FunASR 追踪器一样,为每个 WebSocket 会话保留有界的滚动窗口历史。
|
||||
self.speaker_history: dict[str, dict[str, Any]] = {}
|
||||
self.speaker_last_seen: dict[str, float] = {}
|
||||
self._speaker_cluster_backend: Any | None = None
|
||||
|
|
@ -110,7 +110,7 @@ class AuxiliaryRuntime:
|
|||
raw = os.getenv("AUXILIARY_PRELOAD_KINDS", "")
|
||||
if not raw.strip():
|
||||
return set(DEFAULT_PRELOAD_KINDS)
|
||||
# CAM++ is required; optional model kinds can be added for other endpoints.
|
||||
# CAM++ 是必需模型;其他接口可按需增加可选模型类型。
|
||||
return DEFAULT_PRELOAD_KINDS | {item.strip() for item in raw.split(",") if item.strip()}
|
||||
|
||||
def _load_asset(self, model_id: str, config: dict[str, Any], path: Path) -> Any | None:
|
||||
|
|
@ -119,7 +119,7 @@ class AuxiliaryRuntime:
|
|||
if kind in {"vad", "punctuation"}:
|
||||
from funasr import AutoModel
|
||||
|
||||
# Keep punctuation on CPU by default so CAM++ can retain the GPU.
|
||||
# 标点模型默认在 CPU 上运行,为 CAM++ 预留显存。
|
||||
device = (
|
||||
os.getenv("FUNASR_PUNC_DEVICE", "cpu")
|
||||
if kind == "punctuation"
|
||||
|
|
@ -139,8 +139,8 @@ class AuxiliaryRuntime:
|
|||
|
||||
task = Tasks.speaker_diarization if kind == "diarization" else Tasks.speaker_verification
|
||||
return pipeline(task=task, model=str(path), device=AUXILIARY_DEVICE)
|
||||
# CAM++ 依赖模型和 ForcedAligner 会先确认文件已落盘,后续由各自的专用
|
||||
# 推理路径使用;这里不猜测它们的通用加载方式,避免错误占用显存。
|
||||
# CAM++ 依赖模型会先确认文件已下载,后续交由对应的专用推理流程加载;
|
||||
# 不猜测其他模型的通用加载方式,避免意外占用显存。
|
||||
return None
|
||||
|
||||
def preload(self) -> None:
|
||||
|
|
@ -152,7 +152,7 @@ class AuxiliaryRuntime:
|
|||
kind = str(config.get("kind") or "")
|
||||
path = model_directory(model_id, self.manifest, MODELS_DIR)
|
||||
if kind == "speaker_verification" and os.getenv("CAM_MODEL_PATH"):
|
||||
# The backend launcher passes the selected local CAM++ directory.
|
||||
# 后端启动器会传入选定的本地 CAM++ 目录。
|
||||
path = Path(os.environ["CAM_MODEL_PATH"]).resolve()
|
||||
record: dict[str, Any] = {"path": str(path), "asset_ready": _asset_ready(path, config)}
|
||||
if kind not in preload_kinds:
|
||||
|
|
@ -190,7 +190,7 @@ class AuxiliaryRuntime:
|
|||
failures.append(model_id)
|
||||
self.status[model_id] = record
|
||||
# 启动日志必须包含每个资产的路径、是否完整和底层异常;不能只打印
|
||||
# 一个笼统的“startup failed”,否则远程部署时无法判断缺文件还是版本错误。
|
||||
# 明确说明失败原因,避免只显示笼统的启动失败,便于远程部署时区分文件缺失和版本错误。
|
||||
for model_id, record in self.status.items():
|
||||
print(
|
||||
f"[model] {model_id}: state={record.get('state')}, "
|
||||
|
|
@ -198,7 +198,7 @@ class AuxiliaryRuntime:
|
|||
+ (f", error={record['error']}" if record.get("error") else ""),
|
||||
flush=True,
|
||||
)
|
||||
# The WebSocket service owns the required VAD; this process requires CAM++.
|
||||
# 必需的 VAD 由 WebSocket 服务负责;此进程要求 CAM++ 可用。
|
||||
required_failures = [
|
||||
model_id for model_id in failures
|
||||
if self.assets[model_id].get("kind") == "speaker_verification"
|
||||
|
|
@ -255,7 +255,7 @@ class AuxiliaryRuntime:
|
|||
|
||||
async def vad(self, audio_path: str) -> Any:
|
||||
"""使用临时音频文件执行一次串行化的 VAD 推理。"""
|
||||
# Realtime VAD lives in the WebSocket process; this HTTP endpoint loads on demand.
|
||||
# 实时 VAD 由 WebSocket 进程处理;此 HTTP 接口按需加载模型。
|
||||
try:
|
||||
model = self._find_model("vad")
|
||||
except RuntimeError:
|
||||
|
|
@ -264,12 +264,12 @@ class AuxiliaryRuntime:
|
|||
return await asyncio.to_thread(model.generate, input=audio_path, cache={})
|
||||
|
||||
async def punctuate(self, text: str) -> str:
|
||||
"""Load CT-Transformer lazily and punctuate one completed ASR turn."""
|
||||
"""按需加载 CT-Transformer,为已完成的 ASR 轮次添加标点。"""
|
||||
async with self.inference_lock:
|
||||
try:
|
||||
model = self._find_model("punctuation")
|
||||
except RuntimeError:
|
||||
# Load while holding the lock to prevent duplicate loads across sessions.
|
||||
# 在锁保护下加载,避免多个会话重复加载模型。
|
||||
model = await asyncio.to_thread(self._load_optional_kind, "punctuation")
|
||||
output = await asyncio.to_thread(model.generate, input=text, cache={})
|
||||
if isinstance(output, (list, tuple)) and output:
|
||||
|
|
@ -285,11 +285,11 @@ class AuxiliaryRuntime:
|
|||
try:
|
||||
model = self._find_model("diarization")
|
||||
except RuntimeError:
|
||||
# 完整 diarization 不在核心启动路径,第一次调用接口时才加载。
|
||||
# 完整说话人分离不在核心启动路径中,首次调用对应接口时才加载。
|
||||
model = self._load_optional_kind("diarization")
|
||||
async with self.inference_lock:
|
||||
# ModelScope 的 CAM++ pipeline 以位置参数接收音频路径;使用关键字
|
||||
# input 在不同版本中可能被忽略或直接报参数错误。
|
||||
# ModelScope 的 CAM++ 推理流程以位置参数接收音频路径;使用关键字参数
|
||||
# input 在不同版本中可能被忽略或直接报错。
|
||||
return await asyncio.to_thread(model, audio_path)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -315,21 +315,21 @@ class AuxiliaryRuntime:
|
|||
if result is None:
|
||||
return None
|
||||
|
||||
# ERes2Net pipeline 在 output_emb=True 时返回 {'embs': numpy.ndarray,
|
||||
# 'outputs': ...};部分版本或其它声纹 pipeline 使用 embedding 变体字段。
|
||||
# ERes2Net 推理流程在 output_emb=True 时返回 {'embs': numpy.ndarray,
|
||||
# 'outputs': ...};部分版本或其他声纹推理流程使用 embedding 变体字段。
|
||||
if isinstance(result, Mapping):
|
||||
for key in ("embs", "embedding", "embeddings", "speaker_embedding", "output_embedding"):
|
||||
if key in result:
|
||||
return AuxiliaryRuntime._extract_embedding_value(result[key])
|
||||
return None
|
||||
|
||||
# 某些 ModelScope 版本把结果包装成带 embs/embedding 属性的对象。
|
||||
# 某些 ModelScope 版本会把结果包装成带 embs/embedding 属性的对象。
|
||||
for key in ("embs", "embedding", "embeddings", "speaker_embedding", "output_embedding"):
|
||||
value = getattr(result, key, None)
|
||||
if value is not None:
|
||||
return AuxiliaryRuntime._extract_embedding_value(value)
|
||||
|
||||
# torch.Tensor 不能直接依赖 numpy.asarray 的 object 转换;先显式移到 CPU。
|
||||
# torch.Tensor 不能直接依赖 numpy.asarray 转成普通数组;先显式移到 CPU。
|
||||
detach = getattr(result, "detach", None)
|
||||
if callable(detach):
|
||||
detached = detach()
|
||||
|
|
@ -340,7 +340,7 @@ class AuxiliaryRuntime:
|
|||
if callable(numpy_method):
|
||||
return numpy_method()
|
||||
|
||||
# 单条音频通常返回 [embedding],递归拆开这一层;数值列表则保留为向量。
|
||||
# 单条音频通常返回 [embedding];递归拆开这一层,数值列表则保留为向量。
|
||||
if isinstance(result, (list, tuple)) and len(result) == 1:
|
||||
return AuxiliaryRuntime._extract_embedding_value(result[0])
|
||||
return result
|
||||
|
|
@ -348,14 +348,14 @@ class AuxiliaryRuntime:
|
|||
@staticmethod
|
||||
def _run_embedding_pipeline(model_pipeline: Any, audio_path: str) -> Any:
|
||||
"""调用声纹 pipeline 的公开预处理和 embedding 输出接口。"""
|
||||
# ModelScope 的 ERes2Net pipeline 要求输入为音频路径列表,并通过
|
||||
# output_emb=True 返回 embedding;不能直接把原始 waveform Tensor 喂给
|
||||
# pipeline.model,因为那会跳过采样率、声道和 waveform 预处理。
|
||||
# ModelScope 的 ERes2Net 推理流水线要求输入为音频路径列表,并通过
|
||||
# output_emb=True 返回嵌入向量;不能直接把原始 waveform Tensor 传给
|
||||
# pipeline.model,否则会跳过采样率、声道数和波形预处理。
|
||||
try:
|
||||
result = model_pipeline([audio_path], output_emb=True)
|
||||
except TypeError:
|
||||
# 兼容不支持 output_emb 参数的旧 pipeline:仍然使用 pipeline 自带
|
||||
# preprocess/forward,而不是直接调用内部 model,确保输入格式一致。
|
||||
# 兼容不支持 output_emb 参数的旧推理流水线:仍使用其自带的
|
||||
# preprocess/forward,而不是直接调用内部 model,以确保输入格式一致。
|
||||
preprocess = getattr(model_pipeline, "preprocess", None)
|
||||
forward = getattr(model_pipeline, "forward", None)
|
||||
if not callable(preprocess) or not callable(forward):
|
||||
|
|
@ -414,7 +414,7 @@ class AuxiliaryRuntime:
|
|||
rows = [AuxiliaryRuntime._normalize_embedding(row) for row in array]
|
||||
return rows
|
||||
|
||||
# 一些 ModelScope 版本返回每条音频一个对象,而不是一张二维矩阵。
|
||||
# 某些 ModelScope 版本每条音频返回一个对象,而不是二维矩阵。
|
||||
if isinstance(values, (list, tuple)) and len(values) == expected_count:
|
||||
rows = []
|
||||
for item in values:
|
||||
|
|
@ -470,7 +470,7 @@ class AuxiliaryRuntime:
|
|||
if end <= last_end:
|
||||
break
|
||||
last_end = end
|
||||
# 和 FunASR sv_chunk 一样把尾窗右对齐,确保不足 1.5 秒的末尾窗
|
||||
# 与 FunASR sv_chunk 一样将尾窗右对齐,确保不足 1.5 秒的末尾窗口
|
||||
# 尽量包含完整的新语音,而不是只用较短尾音再补大量零。
|
||||
start = max(0, end - window_samples)
|
||||
chunk = audio[start:end]
|
||||
|
|
@ -498,7 +498,7 @@ class AuxiliaryRuntime:
|
|||
batch_result = model_pipeline(path_batch, output_emb=True)
|
||||
batch_embeddings = self._embedding_rows(batch_result, len(path_batch))
|
||||
except Exception:
|
||||
# 有些旧 pipeline 不接受多文件批次;逐窗调用保持同一预处理路径。
|
||||
# 有些旧推理流水线不接受多文件批次;逐窗调用以保持相同预处理路径。
|
||||
batch_embeddings = [
|
||||
self._normalize_embedding(
|
||||
self._run_embedding_pipeline(model_pipeline, path)
|
||||
|
|
@ -517,7 +517,7 @@ class AuxiliaryRuntime:
|
|||
session_id: str,
|
||||
cluster_centers: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Map FunASR's temporary cluster labels onto stable, capped session IDs."""
|
||||
"""将 FunASR 临时聚类标签映射为稳定且受上限约束的会话编号。"""
|
||||
import numpy as np
|
||||
|
||||
clusters = self.speaker_clusters.setdefault(session_id, [])
|
||||
|
|
@ -556,8 +556,8 @@ class AuxiliaryRuntime:
|
|||
created = True
|
||||
else:
|
||||
if best_cluster is None:
|
||||
# If this turn has more temporary clusters than available IDs,
|
||||
# follow FunASR and fall back to the nearest existing identity.
|
||||
# 如果本轮临时聚类数超过可用编号,
|
||||
# 则按 FunASR 的做法回退到最近的已有身份。
|
||||
best_cluster = max(
|
||||
clusters,
|
||||
key=lambda cluster: float(np.dot(center, cluster["embedding"])),
|
||||
|
|
@ -575,8 +575,8 @@ class AuxiliaryRuntime:
|
|||
best_cluster["count"] = count + 1
|
||||
strategy = "online_embedding_cluster_match"
|
||||
else:
|
||||
# Once the configured identity limit is reached, keep IDs stable
|
||||
# by assigning unmatched clusters to their nearest known center.
|
||||
# 达到配置的身份上限后,保持编号稳定,
|
||||
# 将未匹配的聚类分配给最近的已知中心。
|
||||
strategy = "online_embedding_cluster_limit_fallback"
|
||||
|
||||
used_ids.add(speaker_id)
|
||||
|
|
@ -595,7 +595,7 @@ class AuxiliaryRuntime:
|
|||
self,
|
||||
session_id: str,
|
||||
) -> tuple[list[list[float]], list[dict[str, Any]]]:
|
||||
"""Re-cluster the rolling CAM++ history with FunASR's own backend."""
|
||||
"""使用 FunASR 自带的后端对 CAM++ 滚动历史重新聚类。"""
|
||||
import numpy as np
|
||||
import torch
|
||||
from funasr.models.campplus.cluster_backend import ClusterBackend
|
||||
|
|
@ -613,8 +613,8 @@ class AuxiliaryRuntime:
|
|||
).to("cpu")
|
||||
self._speaker_postprocess = postprocess
|
||||
|
||||
# ClusterBackend produces turn-local labels; postprocess also aligns overlap
|
||||
# boundaries and smooths short speaker runs before stable IDs are assigned.
|
||||
# ClusterBackend 生成本轮标签;后处理还会对齐重叠区间边界,
|
||||
# 并在分配稳定编号前平滑过短的说话人片段。
|
||||
labels = self._speaker_cluster_backend(embeddings, oracle_num=None)
|
||||
labels = np.asarray(labels)
|
||||
chunks = [
|
||||
|
|
@ -638,7 +638,7 @@ class AuxiliaryRuntime:
|
|||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> dict[str, Any]:
|
||||
"""Assign a fallback whole-turn embedding with FunASR's 15-speaker policy."""
|
||||
"""按照 FunASR 的 15 人策略,使用整轮嵌入进行回退分配。"""
|
||||
import numpy as np
|
||||
|
||||
embedding = self._normalize_embedding(embedding)
|
||||
|
|
@ -670,7 +670,7 @@ class AuxiliaryRuntime:
|
|||
confidence = 0.75
|
||||
strategy = "online_embedding_cluster_new"
|
||||
else:
|
||||
# Match FunASR's fallback after the identity limit is reached.
|
||||
# 达到身份上限后,使用与 FunASR 一致的回退策略。
|
||||
speaker_id = int(best_cluster["speaker_id"]) if best_cluster else 0
|
||||
confidence = max(0.0, best_score)
|
||||
strategy = "online_embedding_cluster_limit_fallback"
|
||||
|
|
@ -694,7 +694,7 @@ class AuxiliaryRuntime:
|
|||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Track a completed VAD turn using FunASR's rolling-window clusterer."""
|
||||
"""使用 FunASR 滚动窗口聚类器追踪已完成的 VAD 轮次。"""
|
||||
async with self.inference_lock:
|
||||
now = time.monotonic()
|
||||
for stale_id, seen in list(self.speaker_last_seen.items()):
|
||||
|
|
@ -713,7 +713,7 @@ class AuxiliaryRuntime:
|
|||
}
|
||||
self.speaker_history[session_id] = history
|
||||
|
||||
# Keep absolute timestamps so each turn can be re-clustered against recent turns.
|
||||
# 保留绝对时间戳,以便每轮都能结合近期轮次重新聚类。
|
||||
for window_start, window_end, embedding in windows:
|
||||
history["chunks"].append(
|
||||
(start_time_ms + window_start, start_time_ms + window_end)
|
||||
|
|
@ -756,7 +756,7 @@ class AuxiliaryRuntime:
|
|||
) -> dict[str, Any] | None:
|
||||
"""对一个实时 turn 提取声纹,并更新该 session 的在线聚类中心。"""
|
||||
async with self.inference_lock:
|
||||
# 异常断网时客户端可能来不及 reset,过期状态在下一次请求时回收。
|
||||
# 异常断网时客户端可能来不及发送重置请求;过期状态会在下一次请求时回收。
|
||||
now = time.monotonic()
|
||||
for stale_id, seen in list(self.speaker_last_seen.items()):
|
||||
if now - seen > 1800:
|
||||
|
|
@ -766,7 +766,7 @@ class AuxiliaryRuntime:
|
|||
if embedding is None:
|
||||
return {"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0,
|
||||
"speaker_status": "insufficient_audio", "speaker_reason": "音频不足 800ms,未提取声纹"}
|
||||
# reset 可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。
|
||||
# 重置请求可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。
|
||||
if session_id not in self.speaker_last_seen:
|
||||
return None
|
||||
return self._assign_embedding(
|
||||
|
|
@ -792,7 +792,7 @@ async def health_handler(request: web.Request) -> web.Response:
|
|||
if config.get("kind") == "vad" and model_id in runtime.models),
|
||||
None,
|
||||
)
|
||||
# Readiness here reports the CAM++ model used by the realtime backend.
|
||||
# 此处的就绪状态反映实时后端使用的 CAM++ 模型。
|
||||
ready = speaker_model_id is not None
|
||||
return web.json_response(
|
||||
{
|
||||
|
|
@ -832,7 +832,7 @@ async def vad_handler(request: web.Request) -> web.Response:
|
|||
|
||||
|
||||
async def punctuation_handler(request: web.Request) -> web.Response:
|
||||
"""Return punctuation when the optional local CT-Transformer is available."""
|
||||
"""本地可选 CT-Transformer 可用时返回标点结果。"""
|
||||
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
|
||||
try:
|
||||
payload = await request.json()
|
||||
|
|
@ -847,7 +847,7 @@ async def punctuation_handler(request: web.Request) -> web.Response:
|
|||
punctuated = await runtime.punctuate(text)
|
||||
return web.json_response({"text": punctuated, "available": True})
|
||||
except Exception as exc:
|
||||
# Punctuation is optional; keep ASR usable and expose the reason to the bridge.
|
||||
# 标点模型为可选项;保持 ASR 可用,并将不可用原因告知协议桥接层。
|
||||
print(f"[punctuation] unavailable: {exc}", flush=True)
|
||||
return web.json_response({"text": text, "available": False, "error": str(exc)})
|
||||
|
||||
|
|
@ -887,7 +887,7 @@ def _normalize_diarization_segments(result: Any) -> list[dict[str, Any]]:
|
|||
end_value = _coerce_finite_float(end)
|
||||
if start_value is None or end_value is None:
|
||||
continue
|
||||
# 列表形式是 CAM++ 的秒单位;明确命名为 start_time/end_time 的
|
||||
# 列表形式的 CAM++ 时间单位为秒;明确命名为 start_time/end_time 的字段
|
||||
# 字段按毫秒处理,避免用“超过多少数值”猜单位导致长录音误判。
|
||||
if values_are_milliseconds:
|
||||
start_value /= 1000
|
||||
|
|
@ -927,7 +927,7 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
|
|||
if not session_id:
|
||||
return web.json_response({"error": "multipart field 'session_id' is required"}, status=400)
|
||||
try:
|
||||
# aiohttp 的 MultiDictProxy 值可能是 str、bytes 或 FileField,先收窄
|
||||
# aiohttp 的 MultiDictProxy 值可能是 str、bytes 或 FileField,先转换为明确类型
|
||||
# 为有限浮点数,避免静态检查告警和异常类型值进入声纹服务。
|
||||
start_time_ms = _parse_form_float(form.get("start_time_ms"), "start_time_ms", default=0.0)
|
||||
end_time_ms = _parse_form_float(form.get("end_time_ms"), "end_time_ms", default=start_time_ms)
|
||||
|
|
@ -950,7 +950,7 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
|
|||
"speaker_status": "no_embedding", "speaker_reason": "当前片段未生成可用声纹",
|
||||
})
|
||||
except Exception as exc:
|
||||
# 将模型推理异常返回给 WebSocket 客户端,避免客户端只能看到笼统的 500。
|
||||
# 将模型推理异常返回给 WebSocket 客户端,避免客户端只能看到笼统的 500 错误。
|
||||
print(
|
||||
f"[speaker] resolve failed: session_id={session_id}, "
|
||||
f"start={start_time_ms}, end={end_time_ms}, error={exc}",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""独立 VLLM 部署项目的模型清单与路径解析辅助函数。"""
|
||||
"""FunASR 运行时的模型清单与本地资源路径工具。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -21,7 +21,7 @@ def load_manifest(path: Path = MANIFEST_PATH) -> dict[str, Any]:
|
|||
|
||||
|
||||
def resolve_model_id(model: str | None, manifest: dict[str, Any]) -> str:
|
||||
"""将默认值、短别名或完整模型 ID 解析为一个 ASR 模型。"""
|
||||
"""根据已配置的默认值、别名或模型 ID 查找 ASR 模型。"""
|
||||
models = manifest["models"]
|
||||
requested = (model or "default").strip()
|
||||
if requested == "default":
|
||||
|
|
@ -47,7 +47,7 @@ def model_directory(model_id: str, manifest: dict[str, Any], models_dir: Path) -
|
|||
|
||||
|
||||
def auxiliary_models(manifest: dict[str, Any]) -> dict[str, dict[str, Any]]:
|
||||
"""返回可独立部署的 VAD、说话人和对齐模型资产。"""
|
||||
"""返回已配置的 VAD、说话人、标点及其他辅助模型资源。"""
|
||||
models = manifest.get("auxiliary_models", {})
|
||||
if not isinstance(models, dict):
|
||||
raise ValueError("model_manifest.json 的 auxiliary_models 必须是对象")
|
||||
|
|
@ -58,7 +58,7 @@ def resolve_auxiliary_model_id(
|
|||
manifest: dict[str, Any],
|
||||
kind: str | None = None,
|
||||
) -> str:
|
||||
"""Resolve a configured auxiliary asset by ID, alias, or FunASR alias."""
|
||||
"""根据 ID、别名或 FunASR 别名查找已配置的辅助模型资源。"""
|
||||
assets = auxiliary_models(manifest)
|
||||
requested = (model or "").strip()
|
||||
if requested in assets and (kind is None or assets[requested].get("kind") == kind):
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
"""FunASR realtime browser demo package."""
|
||||
"""FunASR 实时浏览器演示包。"""
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ class AuxiliaryModelService:
|
|||
return decoded
|
||||
|
||||
async def punctuate(self, text: str) -> dict[str, Any]:
|
||||
"""Apply optional FunASR punctuation without making it a startup dependency."""
|
||||
"""应用可选的 FunASR 标点模型,不将其设为启动必需依赖。"""
|
||||
if self._session is None:
|
||||
raise RuntimeError("auxiliary model service is not started")
|
||||
endpoint = self.config.base_url.rstrip("/") + "/v1/punctuation"
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
"""FunASR realtime engine migrated into the demo project.
|
||||
"""FunASR 流式引擎适配器。
|
||||
|
||||
This module keeps the browser-facing project independent from the original
|
||||
Qwen/VLLM service. It follows FunASR's streaming cache lifecycle:
|
||||
one cache per ASR session and is_final=True only for the last chunk.
|
||||
每个会话单独维护流式缓存,并且只将最后一个分块标记为结束块。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -18,7 +16,7 @@ LOGGER = logging.getLogger(__name__)
|
|||
|
||||
@dataclass(frozen=True)
|
||||
class FunASRServiceConfig:
|
||||
"""Configuration for the in-process FunASR models."""
|
||||
"""进程内 FunASR 模型的配置。"""
|
||||
|
||||
model: str = "paraformer-zh-streaming"
|
||||
vad_model: str = "fsmn-vad"
|
||||
|
|
@ -33,7 +31,7 @@ class FunASRServiceConfig:
|
|||
|
||||
@classmethod
|
||||
def from_env(cls) -> "FunASRServiceConfig":
|
||||
"""Read model selection from environment without changing WS fields."""
|
||||
"""从环境变量读取模型选择,不改变 WebSocket 字段。"""
|
||||
chunk_text = os.getenv("FUNASR_CHUNK_SIZE", "0,10,5")
|
||||
try:
|
||||
values = tuple(int(value.strip()) for value in chunk_text.split(","))
|
||||
|
|
@ -61,7 +59,7 @@ class FunASRServiceConfig:
|
|||
|
||||
@dataclass(frozen=True)
|
||||
class FunASRSegment:
|
||||
"""A partial or final event returned by one browser session."""
|
||||
"""浏览器会话返回的中间结果或最终结果事件。"""
|
||||
|
||||
text: str
|
||||
start_time_ms: float
|
||||
|
|
@ -74,7 +72,7 @@ class FunASRSegment:
|
|||
|
||||
|
||||
def _result_text(result: Any) -> str:
|
||||
"""Extract text from the list/dict result shapes used by FunASR."""
|
||||
"""从 FunASR 使用的列表或字典结果结构中提取文本。"""
|
||||
if isinstance(result, list):
|
||||
return _result_text(result[0]) if result else ""
|
||||
if isinstance(result, dict):
|
||||
|
|
@ -86,7 +84,7 @@ def _result_text(result: Any) -> str:
|
|||
|
||||
|
||||
def _vad_events(result: Any) -> list[tuple[float, float]]:
|
||||
"""Normalize FunASR streaming VAD output to start/end milliseconds."""
|
||||
"""将 FunASR 流式 VAD 输出统一转换为毫秒起止时间。"""
|
||||
if isinstance(result, list):
|
||||
result = result[0] if result else {}
|
||||
if isinstance(result, dict):
|
||||
|
|
@ -105,7 +103,7 @@ def _vad_events(result: Any) -> list[tuple[float, float]]:
|
|||
|
||||
|
||||
class FunASRModelService:
|
||||
"""Load FunASR once and create isolated streaming state per browser."""
|
||||
"""只加载一次 FunASR,并为每个浏览器会话创建独立的流状态。"""
|
||||
|
||||
native_partial_supported = True
|
||||
|
||||
|
|
@ -113,12 +111,12 @@ class FunASRModelService:
|
|||
self.config = config or FunASRServiceConfig.from_env()
|
||||
self.asr_model: Any | None = None
|
||||
self.vad_model: Any | None = None
|
||||
# Model objects are shared; each session owns its own cache. Serializing
|
||||
# calls makes the first validation version predictable on one GPU.
|
||||
# 模型对象由会话共享,每个会话单独持有缓存。串行处理调用,
|
||||
# 让首个验证版本在单卡环境中的行为更可预测。
|
||||
self.inference_lock = asyncio.Lock()
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Load streaming ASR and VAD outside the event loop."""
|
||||
"""在事件循环之外加载流式 ASR 和 VAD。"""
|
||||
try:
|
||||
from funasr import AutoModel
|
||||
except ImportError as exc: # pragma: no cover - deployment-only branch
|
||||
|
|
@ -142,18 +140,18 @@ class FunASRModelService:
|
|||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Release references so CUDA memory can be reclaimed on shutdown."""
|
||||
"""释放对象引用,以便关闭时回收 CUDA 显存。"""
|
||||
self.asr_model = None
|
||||
self.vad_model = None
|
||||
|
||||
def create_session(self) -> "FunASRRealtimeSession":
|
||||
"""Create a session with isolated VAD and ASR caches."""
|
||||
"""创建具有独立 VAD 和 ASR 缓存的会话。"""
|
||||
if self.asr_model is None or self.vad_model is None:
|
||||
raise RuntimeError("FunASR model service is not started")
|
||||
return FunASRRealtimeSession(self)
|
||||
|
||||
async def generate_asr(self, audio: Any, status: dict[str, Any]) -> str:
|
||||
"""Run one blocking ASR chunk while preserving its mutable cache."""
|
||||
"""运行一个阻塞式 ASR 分块,同时保留其可变缓存。"""
|
||||
if self.asr_model is None:
|
||||
raise RuntimeError("FunASR ASR model is not loaded")
|
||||
|
||||
|
|
@ -169,7 +167,7 @@ class FunASRModelService:
|
|||
status: dict[str, Any],
|
||||
chunk_ms: int,
|
||||
) -> list[tuple[float, float]]:
|
||||
"""Run one streaming VAD chunk and normalize endpoint events."""
|
||||
"""运行一个流式 VAD 分块,并统一端点事件格式。"""
|
||||
if self.vad_model is None:
|
||||
raise RuntimeError("FunASR VAD model is not loaded")
|
||||
|
||||
|
|
@ -182,7 +180,7 @@ class FunASRModelService:
|
|||
|
||||
|
||||
class _StreamingASRTurn:
|
||||
"""One utterance using FunASR's ordered chunk/cache lifecycle."""
|
||||
"""按照 FunASR 的分块与缓存顺序处理一段话语。"""
|
||||
|
||||
def __init__(self, service: FunASRModelService) -> None:
|
||||
self.service = service
|
||||
|
|
@ -195,7 +193,7 @@ class _StreamingASRTurn:
|
|||
|
||||
@staticmethod
|
||||
def _merge_chunk_text(current: str, chunk: str, previous_chunk: str) -> str:
|
||||
"""Merge chunk text while tolerating wrappers returning cumulative text."""
|
||||
"""合并分块文本,并兼容返回累计文本的封装实现。"""
|
||||
if not chunk or chunk == previous_chunk:
|
||||
return current
|
||||
if current and chunk.startswith(current):
|
||||
|
|
@ -203,12 +201,12 @@ class _StreamingASRTurn:
|
|||
return current + chunk
|
||||
|
||||
async def append(self, pcm_bytes: bytes) -> list[str]:
|
||||
"""Decode complete chunks but hold one chunk for the final flush."""
|
||||
"""解码完整分块,同时暂存一个分块用于最终刷新。"""
|
||||
if not pcm_bytes:
|
||||
return []
|
||||
self.pending.extend(pcm_bytes)
|
||||
chunk_bytes = self.chunk_samples * 2
|
||||
# Hold one full chunk so the real last chunk receives is_final=True.
|
||||
# 暂存一个完整分块,确保真正的最后一块收到 is_final=True。
|
||||
while len(self.pending) >= chunk_bytes * 2:
|
||||
chunk = bytes(self.pending[:chunk_bytes])
|
||||
del self.pending[:chunk_bytes]
|
||||
|
|
@ -234,7 +232,7 @@ class _StreamingASRTurn:
|
|||
return outputs
|
||||
|
||||
async def finish(self) -> str:
|
||||
"""Flush the last buffered chunk and return cumulative text."""
|
||||
"""刷新最后一个缓冲分块并返回累计文本。"""
|
||||
if self.pending:
|
||||
chunk = bytes(self.pending)
|
||||
self.pending.clear()
|
||||
|
|
@ -257,7 +255,7 @@ class _StreamingASRTurn:
|
|||
|
||||
|
||||
class FunASRRealtimeSession:
|
||||
"""FunASR VAD + streaming ASR session used by one browser WebSocket."""
|
||||
"""供单个浏览器 WebSocket 使用的 FunASR VAD 与流式 ASR 会话。"""
|
||||
|
||||
def __init__(self, service: FunASRModelService) -> None:
|
||||
self.service = service
|
||||
|
|
@ -277,13 +275,13 @@ class FunASRRealtimeSession:
|
|||
|
||||
@staticmethod
|
||||
def _to_float32(pcm_bytes: bytes) -> Any:
|
||||
"""Convert browser PCM16 to the float waveform expected by FunASR."""
|
||||
"""将浏览器 PCM16 音频转换为 FunASR 所需的浮点波形。"""
|
||||
import numpy as np
|
||||
|
||||
return np.frombuffer(pcm_bytes, dtype=np.int16).astype(np.float32) / 32768.0
|
||||
|
||||
def _start_segment(self, start_ms: float) -> None:
|
||||
"""Open a turn and retain a short pre-roll for initial phonemes."""
|
||||
"""开启一个轮次,并保留短暂的预录音以覆盖起始音素。"""
|
||||
self.speech_started = True
|
||||
self.segment_start_ms = max(
|
||||
0.0,
|
||||
|
|
@ -295,7 +293,7 @@ class FunASRRealtimeSession:
|
|||
self.pre_roll.clear()
|
||||
|
||||
async def _feed_vad_chunk(self, chunk: bytes) -> list[FunASRSegment]:
|
||||
"""Feed VAD, then stream this audio into the active ASR turn."""
|
||||
"""将音频送入 VAD,再流式传入当前 ASR 轮次。"""
|
||||
self.total_samples += len(chunk) // 2
|
||||
events = await self.service.generate_vad(
|
||||
self._to_float32(chunk),
|
||||
|
|
@ -328,7 +326,7 @@ class FunASRRealtimeSession:
|
|||
return len(self.segment_audio) / (self.sample_rate * 2) * 1000
|
||||
|
||||
def _partial(self, text: str) -> FunASRSegment:
|
||||
"""Create a partial event using the cumulative FunASR text."""
|
||||
"""使用 FunASR 的累计文本创建中间结果事件。"""
|
||||
return FunASRSegment(
|
||||
text=text,
|
||||
start_time_ms=self.segment_start_ms,
|
||||
|
|
@ -340,7 +338,7 @@ class FunASRRealtimeSession:
|
|||
)
|
||||
|
||||
async def _finish_segment(self, reason: str) -> FunASRSegment:
|
||||
"""Flush the ASR cache, then release completed segment audio."""
|
||||
"""刷新 ASR 缓存,然后释放已完成片段的音频。"""
|
||||
text = await self.asr_turn.finish() if self.asr_turn is not None else ""
|
||||
result = FunASRSegment(
|
||||
text=text,
|
||||
|
|
@ -360,7 +358,7 @@ class FunASRRealtimeSession:
|
|||
return result
|
||||
|
||||
async def feed(self, pcm_bytes: bytes) -> list[FunASRSegment]:
|
||||
"""Consume PCM16 and return FunASR partial/final events."""
|
||||
"""接收 PCM16 音频并返回 FunASR 中间或最终结果事件。"""
|
||||
if len(pcm_bytes) % 2:
|
||||
raise ValueError("PCM16 音频必须包含完整的双字节采样")
|
||||
self.vad_buffer.extend(pcm_bytes)
|
||||
|
|
@ -372,7 +370,7 @@ class FunASRRealtimeSession:
|
|||
return results
|
||||
|
||||
async def finish(self) -> list[FunASRSegment]:
|
||||
"""Flush buffered audio and discard all stream caches."""
|
||||
"""刷新缓冲音频并清理所有流式缓存。"""
|
||||
results: list[FunASRSegment] = []
|
||||
if self.vad_buffer:
|
||||
chunk = bytes(self.vad_buffer)
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
"""FunASR realtime WebSocket server adapted from the local FunASR checkout.
|
||||
"""基于本地 FunASR 源码适配的实时 WebSocket 服务端。
|
||||
|
||||
Source: ``runtime/python/websocket/funasr_wss_server.py``. The browser adapter
|
||||
uses FunASR's online WS protocol. Offline ASR, punctuation, and in-process
|
||||
speaker verification are optional; CAM++ runs in the separate auxiliary service.
|
||||
参考源码:``runtime/python/websocket/funasr_wss_server.py``。
|
||||
浏览器适配层使用 FunASR 在线 WebSocket 协议。离线 ASR、标点和进程内说话人验证均为可选项;
|
||||
CAM++ 由独立辅助服务运行。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
|
@ -22,7 +22,7 @@ import torch # 保留不影响
|
|||
|
||||
|
||||
def _bounded_env_float(name: str, default: float, minimum: float, maximum: float) -> float:
|
||||
"""Read and validate a numeric FunASR tuning value from the environment."""
|
||||
"""从环境变量读取并校验 FunASR 数值调优参数。"""
|
||||
raw_value = os.getenv(name, str(default)).strip()
|
||||
try:
|
||||
value = float(raw_value)
|
||||
|
|
@ -34,7 +34,7 @@ def _bounded_env_float(name: str, default: float, minimum: float, maximum: float
|
|||
|
||||
|
||||
def _positive_env_int(name: str, default: int) -> int:
|
||||
"""Read a positive integer environment value with a clear startup error."""
|
||||
"""读取正整数环境变量,参数无效时给出明确的启动错误。"""
|
||||
raw_value = os.getenv(name, str(default)).strip()
|
||||
try:
|
||||
value = int(raw_value)
|
||||
|
|
@ -45,12 +45,12 @@ def _positive_env_int(name: str, default: int) -> int:
|
|||
return value
|
||||
|
||||
|
||||
# An explicit value disables FunASR's duration-based silence schedule for testing.
|
||||
# 显式设置该值后,将关闭 FunASR 按静音时长动态调整阈值的机制,便于测试。
|
||||
VAD_MAX_END_SILENCE_MS = _positive_env_int("FUNASR_VAD_MAX_END_SILENCE_MS", 800)
|
||||
VAD_PARAGRAPH_MAX_END_SILENCE_MS = _positive_env_int(
|
||||
"FUNASR_VAD_PARAGRAPH_MAX_END_SILENCE_MS", 5000
|
||||
)
|
||||
# Higher margins classify weak background noise as non-speech more readily.
|
||||
# 增大裕量可以让较弱的背景噪声更容易被判定为非语音。
|
||||
VAD_SPEECH_NOISE_THRESHOLD = _bounded_env_float(
|
||||
"FUNASR_VAD_SPEECH_NOISE_THRES", 0.6, 0.0, 1.0
|
||||
)
|
||||
|
|
@ -203,7 +203,7 @@ def _safe_int(v, default):
|
|||
return default
|
||||
|
||||
|
||||
# ========= speaker db:加缓存,避免每段都读盘 =========
|
||||
# ========= 说话人数据库:使用缓存,避免每段音频都读取磁盘 =========
|
||||
_SPEAKER_DB_CACHE = {}
|
||||
_SPEAKER_DB_CACHE_TS = 0.0
|
||||
|
||||
|
|
@ -250,9 +250,9 @@ def save_offline_wav_segment_sync(websocket, audio_bytes: bytes, reason: str = "
|
|||
|
||||
fs = int(getattr(websocket, "audio_fs", 16000) or 16000)
|
||||
ch = 1
|
||||
sampwidth = 2 # int16
|
||||
sampwidth = 2 # PCM 使用 16 位采样宽度。
|
||||
|
||||
# int16 对齐
|
||||
# 写入音频帧时保持 int16 采样对齐。
|
||||
if len(audio_bytes) % 2 == 1:
|
||||
audio_bytes = audio_bytes[:-1]
|
||||
if not audio_bytes:
|
||||
|
|
@ -282,7 +282,7 @@ print("model loading")
|
|||
from funasr import AutoModel # noqa
|
||||
|
||||
# ====== 离线 ASR ======
|
||||
# Online deployments leave the offline model unloaded to conserve memory.
|
||||
# 在线部署时不加载离线模型,以节省内存。
|
||||
model_asr = (
|
||||
AutoModel(
|
||||
model=args.asr_model,
|
||||
|
|
@ -297,7 +297,7 @@ model_asr = (
|
|||
else None
|
||||
)
|
||||
|
||||
# streaming asr
|
||||
# 流式 ASR
|
||||
model_asr_streaming = AutoModel(
|
||||
model=args.asr_model_online,
|
||||
model_revision=args.asr_model_online_revision,
|
||||
|
|
@ -308,7 +308,7 @@ model_asr_streaming = AutoModel(
|
|||
disable_log=True,
|
||||
)
|
||||
|
||||
# vad
|
||||
# VAD
|
||||
model_vad = AutoModel(
|
||||
model=args.vad_model,
|
||||
model_revision=args.vad_model_revision,
|
||||
|
|
@ -319,7 +319,7 @@ model_vad = AutoModel(
|
|||
disable_log=True,
|
||||
)
|
||||
|
||||
# punc
|
||||
# 标点模型
|
||||
if args.punc_model != "":
|
||||
model_punc = AutoModel(
|
||||
model=args.punc_model,
|
||||
|
|
@ -333,7 +333,7 @@ if args.punc_model != "":
|
|||
else:
|
||||
model_punc = None
|
||||
|
||||
# CAM++ is loaded by the auxiliary service, avoiding a second GPU copy here.
|
||||
# CAM++ 由辅助服务加载,避免在此进程中重复占用显存。
|
||||
model_sv = (
|
||||
AutoModel(
|
||||
model="iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
|
|
@ -374,7 +374,7 @@ async def run_blocking(fn, *a, sem: asyncio.Semaphore | None = None, **kw):
|
|||
|
||||
|
||||
def _generate_sync(model, audio_or_text, status_dict):
|
||||
# 注意:status_dict 里包含 cache,会被 generate 更新
|
||||
# 注意:status_dict 中包含 cache,模型的 generate 调用会更新该缓存。
|
||||
return model.generate(input=audio_or_text, **status_dict)
|
||||
|
||||
|
||||
|
|
@ -397,7 +397,7 @@ async def clear_websocket():
|
|||
|
||||
|
||||
async def ws_serve(websocket, path=None):
|
||||
# websockets 新版本不会传 path,这里做兼容
|
||||
# 新版 websockets 不会传入 path 参数,这里兼容两种调用方式。
|
||||
if path is None:
|
||||
path = getattr(websocket, "path", None)
|
||||
frames = []
|
||||
|
|
@ -409,7 +409,7 @@ async def ws_serve(websocket, path=None):
|
|||
|
||||
websocket.status_dict_asr = {} # hotword 等
|
||||
websocket.status_dict_asr_online = {"cache": {}, "is_final": False}
|
||||
# Pass the test knob as an explicit FunASR argument so it uses a fixed threshold.
|
||||
# 将测试参数显式传给 FunASR,使其使用固定阈值。
|
||||
websocket.status_dict_vad = {
|
||||
"cache": {},
|
||||
"is_final": False,
|
||||
|
|
@ -533,7 +533,7 @@ async def ws_serve(websocket, path=None):
|
|||
print("bad json message:", e, message[:200])
|
||||
continue
|
||||
|
||||
# Avoid per-message logging during long-running audio sessions.
|
||||
# 长时间音频会话中避免逐条消息写日志。
|
||||
|
||||
end_of_input = False
|
||||
if "is_speaking" in messagejson:
|
||||
|
|
@ -547,7 +547,7 @@ async def ws_serve(websocket, path=None):
|
|||
)
|
||||
|
||||
if "sentence_strategy" in messagejson:
|
||||
# Map the Tencent selector to FunASR's VAD endpoint duration.
|
||||
# 将腾讯前端的选项映射为 FunASR 的 VAD 端点时长。
|
||||
strategy = _safe_int(messagejson["sentence_strategy"], 0)
|
||||
websocket.sentence_strategy = strategy if strategy in (0, 1) else 0
|
||||
websocket.status_dict_vad["max_end_silence_time"] = (
|
||||
|
|
@ -623,7 +623,7 @@ async def ws_serve(websocket, path=None):
|
|||
duration_ms = _pcm_duration_ms(pcm, fs=websocket.audio_fs, ch=1, sampwidth=2)
|
||||
websocket.vad_pre_idx += duration_ms
|
||||
|
||||
# online asr
|
||||
# 在线 ASR
|
||||
frames_asr_online.append(pcm)
|
||||
if websocket.mode in ("2pass", "online"):
|
||||
online_needs_finalization = True
|
||||
|
|
@ -642,7 +642,7 @@ async def ws_serve(websocket, path=None):
|
|||
if speech_start:
|
||||
frames_asr.append(pcm)
|
||||
|
||||
# vad online
|
||||
# 在线 VAD
|
||||
try:
|
||||
speech_start_i, speech_end_i = await async_vad(websocket, pcm)
|
||||
except Exception as e:
|
||||
|
|
@ -650,8 +650,8 @@ async def ws_serve(websocket, path=None):
|
|||
record_error(f"vad inference failed: {e}")
|
||||
speech_start_i, speech_end_i = -1, -1
|
||||
|
||||
# 把 FunASR VAD 的绝对音频时间发给桥接层,用于去掉首尾静音,
|
||||
# 让 CAM++ 滑窗只分析当前 VAD turn,而不是整段会话的缓冲音频。
|
||||
# 将 FunASR VAD 的绝对音频时间发送给桥接层,用于裁掉首尾静音,
|
||||
# 让 CAM++ 滑窗只分析当前 VAD 轮次,而不是整段会话的缓冲音频。
|
||||
if speech_start_i != -1 or speech_end_i != -1:
|
||||
await websocket.send(
|
||||
json.dumps(
|
||||
|
|
@ -674,7 +674,7 @@ async def ws_serve(websocket, path=None):
|
|||
frames_asr = []
|
||||
frames_asr.extend(frames_pre)
|
||||
|
||||
# ========== 3) 2pass:离线阶段触发点 ==========
|
||||
# ========== 3) 2pass:离线阶段的触发位置 ==========
|
||||
if (speech_end_i != -1) or (not websocket.is_speaking):
|
||||
await finalize_online_segment()
|
||||
|
||||
|
|
@ -684,7 +684,7 @@ async def ws_serve(websocket, path=None):
|
|||
audio_in = b"".join(pending_offline_audio)
|
||||
reason = "vad_end" if speech_end_i != -1 else "not_speaking"
|
||||
|
||||
# 保存 wav:放线程池,避免磁盘 IO 卡 loop
|
||||
# 在线程池保存 WAV,避免磁盘 I/O 阻塞事件循环。
|
||||
if websocket.save_offline_segments and audio_in:
|
||||
try:
|
||||
await run_blocking(
|
||||
|
|
@ -745,7 +745,7 @@ async def ws_serve(websocket, path=None):
|
|||
# ===================== 推理:全部改为“线程池 + 限流” =====================
|
||||
|
||||
async def async_vad(websocket, audio_in: bytes):
|
||||
# model_vad.generate 是阻塞的,必须 offload
|
||||
# model_vad.generate 是阻塞操作,必须放到线程池执行。
|
||||
out = await run_blocking(_generate_sync, model_vad, audio_in, websocket.status_dict_vad, sem=SEM_VAD)
|
||||
segments_result = out[0].get("value", [])
|
||||
|
||||
|
|
@ -803,7 +803,7 @@ async def async_asr(websocket, audio_in: bytes):
|
|||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
return
|
||||
|
||||
# 1) ASR(阻塞,线程池执行)
|
||||
# 1) ASR(阻塞操作,在线程池执行)
|
||||
rec_result_list = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr,
|
||||
|
|
@ -837,7 +837,7 @@ async def async_asr(websocket, audio_in: bytes):
|
|||
punc_array = None
|
||||
if model_punc is not None and len(text) > 0:
|
||||
try:
|
||||
# punc 只对文本处理
|
||||
# 标点模型只处理文本,不处理音频。
|
||||
punc_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_punc,
|
||||
|
|
@ -855,7 +855,7 @@ async def async_asr(websocket, audio_in: bytes):
|
|||
except Exception as e:
|
||||
print("punc failed:", e)
|
||||
|
||||
# 4) 构造最终 message
|
||||
# 4) 构造最终消息
|
||||
if len(text) > 0:
|
||||
print("======offline final text:", text)
|
||||
message = {
|
||||
|
|
@ -890,7 +890,7 @@ async def async_asr_online(websocket, audio_in: bytes):
|
|||
if len(audio_in) <= 0 and not websocket.status_dict_asr_online.get("is_final", False):
|
||||
return
|
||||
|
||||
# streaming generate 也是阻塞:线程池执行
|
||||
# 流式 generate 同样是阻塞操作,需要在线程池执行。
|
||||
rec_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr_streaming,
|
||||
|
|
@ -900,15 +900,15 @@ async def async_asr_online(websocket, audio_in: bytes):
|
|||
)
|
||||
rec_result = rec_out[0]
|
||||
|
||||
# 2pass:online 只要 partial,不发 final(final 交给 offline)
|
||||
# 2pass 模式下在线阶段只发送中间结果;最终结果交给离线阶段输出。
|
||||
if websocket.mode == "2pass" and websocket.status_dict_asr_online.get("is_final", False):
|
||||
return
|
||||
|
||||
is_final = bool(
|
||||
websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking)
|
||||
)
|
||||
# 即使最终解码没有新增字符,也必须显式发送 final 事件;桥接层靠它
|
||||
# 结束静音期间的 turn,否则最后一条文本会一直停留在 interim 状态。
|
||||
# 即使最终解码没有新增字符,也必须显式发送 final 事件;桥接层依赖此事件
|
||||
# 结束静音期间的轮次,避免最后一句一直停留在中间态。
|
||||
if rec_result.get("text") or is_final:
|
||||
mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
message = {
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Translate the unchanged Tencent demo protocol to FunASR's native online WS."""
|
||||
"""将未修改的腾讯演示协议转换为 FunASR 原生在线 WebSocket 协议。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -17,7 +17,7 @@ from dotenv import load_dotenv
|
|||
|
||||
try:
|
||||
from websockets.asyncio.client import connect as websocket_connect
|
||||
except ImportError: # websockets before 13 exposes the same client at package root.
|
||||
except ImportError: # websockets 13 之前的版本在包根目录提供相同客户端。
|
||||
from websockets import connect as websocket_connect
|
||||
|
||||
try:
|
||||
|
|
@ -51,7 +51,7 @@ SESSION_REGISTRY_KEY = web.AppKey("sessions", dict)
|
|||
|
||||
|
||||
class IncrementalWavDecoder:
|
||||
"""Read a streamed PCM WAV header and yield its 16 kHz mono PCM payload."""
|
||||
"""读取流式 PCM WAV 文件头,并逐段返回 16 kHz 单声道 PCM 负载。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.buffer = bytearray()
|
||||
|
|
@ -213,7 +213,7 @@ def split_text_by_speaker_segments(
|
|||
|
||||
|
||||
class BrowserSession:
|
||||
"""Own one browser/native WS pair and translate their message contracts."""
|
||||
"""管理浏览器与原生 WebSocket 连接,并转换双方的消息格式。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -252,7 +252,7 @@ class BrowserSession:
|
|||
self.native_error: str | None = None
|
||||
|
||||
async def emit(self, payload: dict[str, Any]) -> None:
|
||||
"""Serialize browser writes because ASR and CAM++ finish independently."""
|
||||
"""ASR 和 CAM++ 的完成时机不同,因此需要串行写入浏览器连接。"""
|
||||
async with self.send_lock:
|
||||
if not self.browser_ws.closed:
|
||||
await self.browser_ws.send_json(payload)
|
||||
|
|
@ -273,7 +273,7 @@ class BrowserSession:
|
|||
"sentence_type": 1 if final else 0,
|
||||
"start_time": round(self.turn_start_ms if start_time_ms is None else start_time_ms),
|
||||
"end_time": round(self.total_audio_ms if end_time_ms is None else end_time_ms),
|
||||
# The unchanged Tencent UI uses speaker_id to choose its speaker bubble.
|
||||
# 未修改的腾讯界面使用 speaker_id 选择对应的说话人气泡。
|
||||
"speaker_id": int(speaker.get("speaker_id", -1)),
|
||||
"speaker_name": str(speaker.get("speaker_name") or ""),
|
||||
"speaker_confidence": float(speaker.get("speaker_confidence") or 0),
|
||||
|
|
@ -311,7 +311,7 @@ class BrowserSession:
|
|||
self.total_audio_ms += len(frame) / PCM_BYTES_PER_MS
|
||||
self.turn_audio.extend(frame)
|
||||
if len(self.turn_audio) > MAX_SPEAKER_AUDIO_BYTES:
|
||||
# Bound per-turn RAM even if VAD never reports an endpoint.
|
||||
# 即使 VAD 始终没有报告端点,也要限制每轮音频的内存占用。
|
||||
trim = len(self.turn_audio) - MAX_SPEAKER_AUDIO_BYTES
|
||||
del self.turn_audio[:trim]
|
||||
self.turn_start_ms += trim / PCM_BYTES_PER_MS
|
||||
|
|
@ -341,7 +341,7 @@ class BrowserSession:
|
|||
)
|
||||
|
||||
async def read_native(self, native_ws: Any) -> None:
|
||||
"""Consume native FunASR events and keep its per-utterance partial cache."""
|
||||
"""接收 FunASR 原生事件,并维护每段话语的中间结果缓存。"""
|
||||
try:
|
||||
while True:
|
||||
raw = await native_ws.recv()
|
||||
|
|
@ -354,7 +354,7 @@ class BrowserSession:
|
|||
continue
|
||||
text = str(message.get("text") or "")
|
||||
if text:
|
||||
# FunASR online sends the newly decoded text for each chunk.
|
||||
# FunASR 在线模式会为每个分块发送新解码出的文本。
|
||||
self.turn_text += text
|
||||
if message.get("is_final"):
|
||||
final_text = self.turn_text.strip()
|
||||
|
|
@ -363,8 +363,8 @@ class BrowserSession:
|
|||
end_ms = start_ms + len(audio) / PCM_BYTES_PER_MS
|
||||
turn_sentence_id = self.sentence_id
|
||||
if final_text:
|
||||
# 先结束前端 interim 气泡;标点和 CAM++ 在独立 worker 中完成,
|
||||
# 不阻塞 native WS 继续读取后续音频帧。
|
||||
# 先结束前端中间结果气泡;标点和 CAM++ 在独立工作线程中完成,
|
||||
# 不阻塞原生 WebSocket 继续读取后续音频帧。
|
||||
await self.emit_sentence(
|
||||
final_text,
|
||||
final=True,
|
||||
|
|
@ -381,7 +381,7 @@ class BrowserSession:
|
|||
end_time_ms=end_ms,
|
||||
)
|
||||
)
|
||||
# 为同一个 VAD turn 内可能拆出的多个气泡预留独立 ID。
|
||||
# 为同一个 VAD 轮次内可能拆出的多个气泡预留独立 ID。
|
||||
self.sentence_id += TURN_SENTENCE_ID_STRIDE
|
||||
self.turn_text = ""
|
||||
self.turn_audio.clear()
|
||||
|
|
@ -419,8 +419,8 @@ class BrowserSession:
|
|||
|
||||
subsegments: list[dict[str, Any]] = []
|
||||
if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
|
||||
# required CAM++ 按 FunASR 1.5s/0.75s 滑窗识别同一 VAD turn 内的
|
||||
# 多人切换;若窗长不足或接口暂不可用,再退回整段声纹验证。
|
||||
# 必需的 CAM++ 按 FunASR 的 1.5 秒窗口和 0.75 秒步长,
|
||||
# 在同一个 VAD 轮次内识别多个说话人片段。
|
||||
track = getattr(self.auxiliary, "track_speakers", None)
|
||||
if callable(track):
|
||||
try:
|
||||
|
|
@ -485,7 +485,7 @@ class BrowserSession:
|
|||
|
||||
|
||||
async def config_handler(_: web.Request) -> web.Response:
|
||||
"""Expose a small readiness response for the launcher and diagnostics."""
|
||||
"""提供简洁的就绪状态响应,供启动器和诊断使用。"""
|
||||
return web.json_response(
|
||||
{
|
||||
"engine": "funasr-native-online-ws",
|
||||
|
|
@ -510,7 +510,7 @@ async def stop_handler(request: web.Request) -> web.Response:
|
|||
|
||||
|
||||
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
||||
"""Bridge Tencent's browser messages to FunASR's native realtime protocol."""
|
||||
"""将腾讯前端消息桥接到 FunASR 原生实时协议。"""
|
||||
browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30)
|
||||
await browser_ws.prepare(request)
|
||||
session: BrowserSession | None = None
|
||||
|
|
@ -550,7 +550,7 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
|||
await session.emit({"type": "voice_id", "voice_id": voice_id})
|
||||
await session.emit({"type": "start"})
|
||||
|
||||
# This is FunASR's native WSS message contract; PCM frames follow at 60 ms.
|
||||
# 这是 FunASR 原生 WSS 消息格式;之后每 60 毫秒发送一帧 PCM。
|
||||
async with websocket_connect(
|
||||
NATIVE_WS_URL,
|
||||
subprotocols=["binary"],
|
||||
|
|
@ -620,7 +620,7 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
|||
|
||||
if session.native_error is None and not browser_ws.closed:
|
||||
await session.finish_audio(native_ws)
|
||||
# FunASR flushes its online cache and acknowledges only after final output.
|
||||
# FunASR 会刷新在线缓存,并在输出最终结果后才返回确认。
|
||||
await native_ws.send(
|
||||
json.dumps({"is_speaking": False, "is_end": True}, ensure_ascii=False)
|
||||
)
|
||||
|
|
@ -671,7 +671,7 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
|||
|
||||
|
||||
async def create_app() -> web.Application:
|
||||
"""Create a light protocol bridge; model inference belongs to native FunASR."""
|
||||
"""创建轻量协议桥接层;模型推理由 FunASR 原生服务负责。"""
|
||||
app = web.Application()
|
||||
app[AUXILIARY_KEY] = AuxiliaryModelService(
|
||||
AuxiliaryServiceConfig(
|
||||
|
|
|
|||
|
|
@ -1,256 +0,0 @@
|
|||
"""独立实时 Demo 使用的 OpenAI 兼容 VLLM 服务适配器。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import wave
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from aiohttp import ClientSession, ClientTimeout, FormData, WSMsgType
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelServiceConfig:
|
||||
"""一个独立 VLLM 端点所需的连接配置。"""
|
||||
|
||||
base_url: str = "http://127.0.0.1:9950/v1"
|
||||
model: str = "Qwen/Qwen3-ASR-0.6B"
|
||||
api_key: str = "EMPTY"
|
||||
timeout_seconds: float = 45.0
|
||||
realtime_enabled: bool = True
|
||||
|
||||
|
||||
def pcm16_to_wav(pcm_bytes: bytes, sample_rate: int = 16000) -> bytes:
|
||||
"""将浏览器发送的 PCM16 单声道数据封装为 VLLM 可识别的 WAV 请求。"""
|
||||
output = io.BytesIO()
|
||||
with wave.open(output, "wb") as wav_file:
|
||||
wav_file.setnchannels(1)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(sample_rate)
|
||||
wav_file.writeframes(pcm_bytes)
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
def wav_to_pcm16(audio_bytes: bytes) -> bytes:
|
||||
"""从 WAV 缓冲区提取 PCM 帧,并兼容尚未完整的中间音频数据。"""
|
||||
try:
|
||||
with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file:
|
||||
return wav_file.readframes(wav_file.getnframes())
|
||||
except (EOFError, wave.Error):
|
||||
if audio_bytes[:4] == b"RIFF" and audio_bytes[8:12] == b"WAVE" and len(audio_bytes) > 44:
|
||||
return audio_bytes[44:]
|
||||
return audio_bytes
|
||||
|
||||
|
||||
def prepare_audio_request(
|
||||
audio_bytes: bytes,
|
||||
source: str,
|
||||
file_name: str,
|
||||
partial: bool,
|
||||
) -> tuple[bytes, str, str] | None:
|
||||
"""将麦克风、PCM 或 WAV 数据转换为 WAV;压缩格式的中间片段延迟到最终帧处理。"""
|
||||
suffix = Path(file_name).suffix.lower()
|
||||
if source == "mic" or suffix in {".pcm", ".wav"}:
|
||||
pcm_bytes = wav_to_pcm16(audio_bytes) if suffix == ".wav" else audio_bytes
|
||||
return pcm16_to_wav(pcm_bytes), "audio.wav", "audio/wav"
|
||||
if partial:
|
||||
# MP3/M4A/OGG 的不断增长前缀通常不是完整容器,不能安全解码,因此只在
|
||||
# 最终阶段提交压缩文件,避免中间请求产生随机解码错误。
|
||||
return None
|
||||
content_type = {
|
||||
".mp3": "audio/mpeg",
|
||||
".m4a": "audio/mp4",
|
||||
".ogg": "audio/ogg",
|
||||
".opus": "audio/ogg",
|
||||
}.get(suffix, "application/octet-stream")
|
||||
return audio_bytes, Path(file_name).name or "audio.bin", content_type
|
||||
|
||||
def realtime_ws_url(base_url: str) -> str:
|
||||
"""Convert the configured OpenAI-compatible base URL to vLLM's realtime URL."""
|
||||
parsed = urlsplit(base_url.rstrip("/"))
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
||||
raise ValueError(f"invalid VLLM base URL: {base_url}")
|
||||
scheme = "wss" if parsed.scheme == "https" else "ws"
|
||||
path = parsed.path.rstrip("/") + "/realtime"
|
||||
return urlunsplit((scheme, parsed.netloc, path, "", ""))
|
||||
|
||||
|
||||
class VLLMRealtimeStream:
|
||||
"""One vLLM realtime stream, isolated from the shared HTTP client session."""
|
||||
|
||||
def __init__(self, websocket: Any, timeout_seconds: float) -> None:
|
||||
self._websocket = websocket
|
||||
self._timeout_seconds = timeout_seconds
|
||||
self._latest_text = ""
|
||||
self._error: Exception | None = None
|
||||
self._done = asyncio.Event()
|
||||
self._reader = asyncio.create_task(self._read_messages())
|
||||
|
||||
async def _read_messages(self) -> None:
|
||||
"""Collect model deltas continuously so audio ingestion never waits for a snapshot."""
|
||||
try:
|
||||
async for message in self._websocket:
|
||||
if message.type == WSMsgType.TEXT:
|
||||
try:
|
||||
payload = json.loads(message.data)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
event_type = payload.get("type")
|
||||
if event_type == "transcription.delta":
|
||||
delta = str(payload.get("delta") or "")
|
||||
if delta:
|
||||
self._latest_text += delta
|
||||
elif payload.get("text") is not None:
|
||||
self._latest_text = str(payload["text"])
|
||||
elif event_type == "transcription.done":
|
||||
self._latest_text = str(
|
||||
payload.get("text") or payload.get("transcript") or self._latest_text
|
||||
).strip()
|
||||
self._done.set()
|
||||
elif event_type == "error":
|
||||
detail = payload.get("error") or payload.get("message") or "unknown realtime error"
|
||||
self._error = RuntimeError(str(detail))
|
||||
self._done.set()
|
||||
elif message.type == WSMsgType.ERROR:
|
||||
self._error = self._websocket.exception() or RuntimeError("VLLM realtime WebSocket failed")
|
||||
self._done.set()
|
||||
return
|
||||
elif message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.CLOSING}:
|
||||
if not self._done.is_set():
|
||||
self._error = RuntimeError("VLLM realtime WebSocket closed before transcription.done")
|
||||
self._done.set()
|
||||
return
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
self._error = exc
|
||||
self._done.set()
|
||||
|
||||
def latest_text(self) -> str:
|
||||
"""Return the newest model text already received by the reader task."""
|
||||
return self._latest_text.strip()
|
||||
|
||||
async def append_audio(self, pcm_bytes: bytes) -> None:
|
||||
"""Push one raw 16 kHz mono PCM16 block without re-uploading old audio."""
|
||||
if not pcm_bytes:
|
||||
return
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
await self._websocket.send_json(
|
||||
{
|
||||
"type": "input_audio_buffer.append",
|
||||
"audio": base64.b64encode(pcm_bytes).decode("ascii"),
|
||||
}
|
||||
)
|
||||
|
||||
async def finish(self) -> str:
|
||||
"""Commit the current model turn and return the final realtime transcription."""
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
await self._websocket.send_json({"type": "input_audio_buffer.commit", "final": True})
|
||||
try:
|
||||
await asyncio.wait_for(self._done.wait(), timeout=self._timeout_seconds)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise TimeoutError("VLLM realtime transcription timed out") from exc
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
return self._latest_text.strip()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Stop the reader and release the model-side WebSocket."""
|
||||
if not self._reader.done():
|
||||
self._reader.cancel()
|
||||
await asyncio.gather(self._reader, return_exceptions=True)
|
||||
if not self._websocket.closed:
|
||||
await self._websocket.close()
|
||||
|
||||
|
||||
class VLLMTranscriptionService:
|
||||
"""只调用独立项目提供的 VLLM HTTP 接口,不导入原项目应用代码。"""
|
||||
|
||||
@property
|
||||
def native_partial_supported(self) -> bool:
|
||||
"""Report whether this adapter is configured to use vLLM realtime."""
|
||||
return self.config.realtime_enabled
|
||||
|
||||
def __init__(self, config: ModelServiceConfig) -> None:
|
||||
self.config = config
|
||||
self._session: ClientSession | None = None
|
||||
|
||||
async def start(self) -> None:
|
||||
"""创建可复用的 HTTP 会话,供所有中间和最终转写请求共享。"""
|
||||
self._session = ClientSession(timeout=ClientTimeout(total=self.config.timeout_seconds))
|
||||
|
||||
async def close(self) -> None:
|
||||
"""本地 Demo 退出时释放可复用的 HTTP 会话和底层连接。"""
|
||||
if self._session is not None:
|
||||
await self._session.close()
|
||||
self._session = None
|
||||
|
||||
async def open_realtime_stream(self) -> VLLMRealtimeStream:
|
||||
"""Open a model-native stream; the caller owns and closes the returned turn."""
|
||||
if not self.config.realtime_enabled:
|
||||
raise RuntimeError("VLLM realtime streaming is disabled")
|
||||
if self._session is None:
|
||||
raise RuntimeError("model service is not started")
|
||||
endpoint = realtime_ws_url(self.config.base_url)
|
||||
headers = {"Authorization": f"Bearer {self.config.api_key}"}
|
||||
connect_timeout = min(5.0, self.config.timeout_seconds)
|
||||
websocket = await self._session.ws_connect(
|
||||
endpoint,
|
||||
headers=headers,
|
||||
timeout=connect_timeout,
|
||||
heartbeat=20,
|
||||
)
|
||||
try:
|
||||
created = await asyncio.wait_for(websocket.receive(), timeout=connect_timeout)
|
||||
if created.type == WSMsgType.TEXT:
|
||||
payload = json.loads(created.data)
|
||||
if payload.get("type") == "error":
|
||||
raise RuntimeError(
|
||||
str(payload.get("error") or payload.get("message") or "VLLM realtime error")
|
||||
)
|
||||
await websocket.send_json({"type": "session.update", "model": self.config.model})
|
||||
return VLLMRealtimeStream(websocket, self.config.timeout_seconds)
|
||||
except Exception:
|
||||
await websocket.close()
|
||||
raise
|
||||
|
||||
async def transcribe(
|
||||
self,
|
||||
audio_bytes: bytes,
|
||||
source: str,
|
||||
file_name: str,
|
||||
partial: bool,
|
||||
) -> str | None:
|
||||
"""提交一次音频快照并返回文本;返回 None 表示当前格式不支持中间转写。"""
|
||||
prepared = prepare_audio_request(audio_bytes, source, file_name, partial)
|
||||
if prepared is None:
|
||||
return None
|
||||
payload, upload_name, content_type = prepared
|
||||
if self._session is None:
|
||||
raise RuntimeError("model service is not started")
|
||||
|
||||
form = FormData()
|
||||
form.add_field("file", payload, filename=upload_name, content_type=content_type)
|
||||
form.add_field("model", self.config.model)
|
||||
form.add_field("response_format", "json")
|
||||
headers = {"Authorization": f"Bearer {self.config.api_key}"}
|
||||
endpoint = self.config.base_url.rstrip("/") + "/audio/transcriptions"
|
||||
async with self._session.post(endpoint, data=form, headers=headers) as response:
|
||||
body = await response.text()
|
||||
if response.status >= 400:
|
||||
raise RuntimeError(f"VLLM transcription failed ({response.status}): {body[:500]}")
|
||||
try:
|
||||
decoded: Any = await response.json(content_type=None)
|
||||
except ValueError:
|
||||
return body.strip()
|
||||
if isinstance(decoded, dict):
|
||||
return str(decoded.get("text") or decoded.get("transcript") or "").strip()
|
||||
return str(decoded).strip()
|
||||
|
|
@ -1,855 +0,0 @@
|
|||
"""面向浏览器的独立 Qwen3-ASR VLLM WebSocket 编排服务。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import math
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import webbrowser
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
from uuid import uuid4
|
||||
|
||||
from aiohttp import WSMsgType, web
|
||||
from dotenv import load_dotenv
|
||||
|
||||
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||||
from model_service import ModelServiceConfig, VLLMTranscriptionService
|
||||
from speaker_assembler import SegmentAssembler
|
||||
|
||||
|
||||
# 与部署启动器读取同一配置;外部环境变量优先于 demo/.env。
|
||||
DEPLOY_ROOT = Path(__file__).resolve().parents[2]
|
||||
load_dotenv(DEPLOY_ROOT / ".env")
|
||||
|
||||
# 监听所有网卡,允许同一局域网内的浏览器访问服务器上的 Demo;端口集中在代码
|
||||
# 变量中维护,便于服务器部署时直接修改并保持页面和 WebSocket 使用一致端口。
|
||||
WEB_HOST = "0.0.0.0"
|
||||
WEB_PORT = 8082
|
||||
WEB_DISPLAY_HOST = os.getenv("WEB_DISPLAY_HOST", "127.0.0.1")
|
||||
DEFAULT_MODEL_SERVICE_URL = f"http://127.0.0.1:{os.getenv('VLLM_PORT', '9950')}/v1"
|
||||
PARTIAL_BYTES_PER_SECOND = 16000 * 2
|
||||
VAD_FRAME_BYTES = 640
|
||||
VAD_FRAME_MS = 20
|
||||
VAD_SILENCE_MS = 800
|
||||
PARAGRAPH_SILENCE_MS = 1400
|
||||
VAD_RMS_THRESHOLD = 450
|
||||
MIN_SPEAKER_VOICE_MS = 800
|
||||
# 说话人分离开启时,用比普通 VAD 更短的静音作为候选话轮边界。
|
||||
# 这个边界只负责把 A→B 的短交接停顿拆开,不调用声纹模型,因此不会
|
||||
# 把窗口推理延迟叠加到音频帧处理路径;同一说话人的拆分片段仍由聚类合并。
|
||||
SPEAKER_GAP_MS = max(20, int(os.getenv("SPEAKER_GAP_MS", "400")))
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class EndOfStream:
|
||||
"""带明确类型的队列结束标记,用于区分控制信号和真实音频字节。"""
|
||||
|
||||
|
||||
EOF = EndOfStream()
|
||||
MODEL_SERVICE_KEY = web.AppKey("model_service", VLLMTranscriptionService)
|
||||
AUXILIARY_SERVICE_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SpeakerJob:
|
||||
"""等待辅助服务处理的单个已完成 turn;只保存该 turn 的 PCM 音频。"""
|
||||
|
||||
sentence_id: int
|
||||
audio: bytes
|
||||
start_time_ms: float
|
||||
end_time_ms: float
|
||||
voiced_ms: float = 0.0
|
||||
|
||||
|
||||
def validate_model_service_url(value: str) -> str:
|
||||
"""只接受用户输入的 HTTP(S) VLLM 地址,并拒绝附带认证和查询参数的地址。"""
|
||||
candidate = value.strip().rstrip("/")
|
||||
parsed = urlparse(candidate)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
|
||||
raise ValueError("VLLM 地址必须是完整的 http:// 或 https:// URL")
|
||||
if parsed.username or parsed.password or parsed.query or parsed.fragment:
|
||||
raise ValueError("VLLM 地址不能包含账号、密码、查询参数或片段")
|
||||
return candidate
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionMetrics:
|
||||
"""记录 WebSocket 会话耗时和结果修订次数,并在会话结束后展示。"""
|
||||
|
||||
started_at: float
|
||||
audio_bytes: int = 0
|
||||
input_chunks: int = 0
|
||||
partial_count: int = 0
|
||||
partial_revisions: int = 0
|
||||
first_partial_ms: float | None = None
|
||||
final_ms: float | None = None
|
||||
|
||||
def snapshot(self) -> dict[str, Any]:
|
||||
"""返回可安全序列化为 JSON 的指标,耗时均相对于会话开始时间计算。"""
|
||||
now = time.perf_counter()
|
||||
return {
|
||||
"audio_bytes": self.audio_bytes,
|
||||
"input_chunks": self.input_chunks,
|
||||
"partial_count": self.partial_count,
|
||||
"partial_revisions": self.partial_revisions,
|
||||
"first_partial_ms": self.first_partial_ms,
|
||||
"final_ms": self.final_ms,
|
||||
"elapsed_ms": round((now - self.started_at) * 1000, 1),
|
||||
}
|
||||
|
||||
|
||||
class RealtimeSession:
|
||||
"""串行处理音频快照,并通过同一个 WebSocket 有序推送状态更新。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ws: web.WebSocketResponse,
|
||||
model_service: VLLMTranscriptionService,
|
||||
auxiliary_service: AuxiliaryModelService | None,
|
||||
start: dict[str, Any],
|
||||
) -> None:
|
||||
self.ws = ws
|
||||
self.model_service = model_service
|
||||
self.auxiliary_service = auxiliary_service
|
||||
self.start = start
|
||||
# 聚类状态只能属于当前连接,客户端复用 ID 不能串入另一会话的声纹池。
|
||||
self.session_id = uuid4().hex
|
||||
self.send_lock = asyncio.Lock()
|
||||
self.state_revision = 0
|
||||
self.audio_queue: asyncio.Queue[bytes | EndOfStream] = asyncio.Queue(maxsize=256)
|
||||
self.speaker_queue: asyncio.Queue[SpeakerJob | EndOfStream] = asyncio.Queue(maxsize=64)
|
||||
self.assembler = SegmentAssembler()
|
||||
self.metrics = SessionMetrics(time.perf_counter())
|
||||
self.source = str(start.get("source") or "mic")
|
||||
self.file_name = str(start.get("file_name") or "audio.wav")
|
||||
self.windowed_partial = self.source == "mic" or Path(self.file_name).suffix.lower() in {".pcm", ".wav"}
|
||||
self.sentence_strategy = int(start.get("sentence_strategy") or 0)
|
||||
self.silence_limit_ms = PARAGRAPH_SILENCE_MS if self.sentence_strategy == 1 else VAD_SILENCE_MS
|
||||
self.partial_interval_ms = max(300, int(start.get("partial_interval_ms") or 1200))
|
||||
self.max_segment_sec = max(2.0, float(start.get("max_segment_sec") or 12.0))
|
||||
self.merge_adjacent = self._parse_flag(start.get("display_merge"), True)
|
||||
self.enable_native_partial = self._parse_flag(start.get("enable_native_partial_stream"), True)
|
||||
self.native_stream: Any | None = None
|
||||
self.native_stream_disabled = False
|
||||
self.native_sent_bytes = 0
|
||||
self.segment_id = 0
|
||||
self.segment_audio = bytearray()
|
||||
self.segment_start_ms = 0.0
|
||||
self.vad_buffer = bytearray()
|
||||
self.processed_audio_bytes = 0
|
||||
self.silence_ms = 0
|
||||
self.in_speech = False
|
||||
self.voiced_ms = 0.0
|
||||
self.pre_roll = bytearray()
|
||||
self.wav_header_buffer = bytearray()
|
||||
self.wav_payload_started = self.source != "file" or Path(self.file_name).suffix.lower() != ".wav"
|
||||
self.wav_riff_read = False
|
||||
self.wav_format_valid = False
|
||||
self.wav_data_remaining: int | None = None
|
||||
self.speaker_warning_sent = False
|
||||
self.speaker_enabled = self._parse_flag(start.get("speaker_diarization"), True)
|
||||
# 没有辅助服务时不提前切段,避免服务不可用时把一段语音拆成许多未知片段。
|
||||
self.speaker_gap_enabled = self.speaker_enabled and auxiliary_service is not None
|
||||
raw_speaker_gap = start.get("speaker_gap_ms")
|
||||
if raw_speaker_gap is None:
|
||||
self.speaker_gap_ms = SPEAKER_GAP_MS
|
||||
else:
|
||||
try:
|
||||
self.speaker_gap_ms = max(20, int(float(raw_speaker_gap)))
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
self.speaker_gap_ms = SPEAKER_GAP_MS
|
||||
self.input_stopped = False
|
||||
|
||||
@staticmethod
|
||||
def _parse_flag(value: Any, default: bool) -> bool:
|
||||
"""兼容前端传来的 0/1、布尔值和字符串开关,避免字符串 0 被误判为真。"""
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() not in {"", "0", "false", "no", "off"}
|
||||
return bool(value)
|
||||
|
||||
async def emit(self, payload: dict[str, Any]) -> None:
|
||||
"""在连接仍然有效时发送一条有序事件,避免向已关闭连接写入数据。"""
|
||||
async with self.send_lock:
|
||||
if not self.ws.closed:
|
||||
await self.ws.send_json(payload)
|
||||
|
||||
async def emit_state(self, sentence: dict[str, Any] | None = None) -> None:
|
||||
"""每次状态更新后同时发送原始状态和重新计算的展示快照。"""
|
||||
# 在首次 await 前冻结快照,音频 worker 与 speaker worker 不会混用两版状态。
|
||||
self.state_revision += 1
|
||||
state = {
|
||||
"type": "display_state", "revision": self.state_revision,
|
||||
"raw_segments": self.assembler.raw_snapshot(),
|
||||
"display_blocks": self.assembler.display_blocks(self.merge_adjacent),
|
||||
"metrics": self.metrics.snapshot(),
|
||||
}
|
||||
if sentence is not None:
|
||||
await self.emit({"type": "sentences", "sentences": [sentence], "metrics": self.metrics.snapshot()})
|
||||
await self.emit(state)
|
||||
|
||||
async def warn_speaker(self, message: str) -> None:
|
||||
"""只发送一次说话人服务告警,避免辅助服务异常时刷屏。"""
|
||||
if self.speaker_warning_sent:
|
||||
return
|
||||
await self.emit(
|
||||
{
|
||||
"type": "speaker_warning",
|
||||
"session_id": self.session_id,
|
||||
"speaker_service_url": getattr(getattr(self.auxiliary_service, "config", None), "base_url", None),
|
||||
"message": message,
|
||||
}
|
||||
)
|
||||
self.speaker_warning_sent = True
|
||||
|
||||
def _duration_ms(self) -> float:
|
||||
"""根据 16 kHz PCM 字节数计算时长,不依赖前端可能漂移的时间戳。"""
|
||||
return len(self.segment_audio) / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
|
||||
async def _transcribe(self, partial: bool) -> str | None:
|
||||
"""通过 VLLM 适配器转写当前逻辑片段,并保留中间/最终请求的统一入口。"""
|
||||
return await self.model_service.transcribe(
|
||||
bytes(self.segment_audio),
|
||||
"mic",
|
||||
"turn.pcm",
|
||||
partial=partial,
|
||||
)
|
||||
|
||||
async def _start_native_stream(self) -> None:
|
||||
"""Start one model stream for the current VAD turn; HTTP remains the safe fallback."""
|
||||
if not self.enable_native_partial or self.native_stream_disabled or self.native_stream is not None:
|
||||
return
|
||||
opener = getattr(self.model_service, "open_realtime_stream", None)
|
||||
if not callable(opener):
|
||||
self.native_stream_disabled = True
|
||||
return
|
||||
try:
|
||||
self.native_stream = await opener()
|
||||
self.native_sent_bytes = 0
|
||||
except Exception as exc:
|
||||
self.native_stream_disabled = True
|
||||
LOGGER.warning("native ASR stream unavailable; falling back to HTTP: %s", exc)
|
||||
|
||||
async def _feed_native_audio(self) -> None:
|
||||
"""Send only new PCM bytes so a long turn is never re-uploaded as a growing snapshot."""
|
||||
if self.native_stream is None:
|
||||
return
|
||||
pending = bytes(self.segment_audio[self.native_sent_bytes:])
|
||||
if not pending:
|
||||
return
|
||||
try:
|
||||
await self.native_stream.append_audio(pending)
|
||||
self.native_sent_bytes = len(self.segment_audio)
|
||||
except Exception as exc:
|
||||
await self._close_native_stream()
|
||||
self.native_stream_disabled = True
|
||||
LOGGER.warning("native ASR stream failed; falling back to HTTP: %s", exc)
|
||||
|
||||
async def _partial_text(self) -> str | None:
|
||||
"""Read the latest native delta; use the old HTTP path only when native streaming is unavailable."""
|
||||
if self.native_stream is not None:
|
||||
return self.native_stream.latest_text()
|
||||
return await self._transcribe(partial=True)
|
||||
|
||||
async def _finish_transcription(self) -> str | None:
|
||||
"""Commit the native turn once, then fall back to one final HTTP request on failure."""
|
||||
stream = self.native_stream
|
||||
self.native_stream = None
|
||||
self.native_sent_bytes = 0
|
||||
if stream is None:
|
||||
return await self._transcribe(partial=False)
|
||||
try:
|
||||
return await stream.finish()
|
||||
except Exception as exc:
|
||||
self.native_stream_disabled = True
|
||||
LOGGER.warning("native ASR finalization failed; falling back to HTTP: %s", exc)
|
||||
return await self._transcribe(partial=False)
|
||||
finally:
|
||||
await stream.close()
|
||||
|
||||
async def _close_native_stream(self) -> None:
|
||||
"""Release an unfinished native turn when the browser disconnects or aborts."""
|
||||
stream = self.native_stream
|
||||
self.native_stream = None
|
||||
self.native_sent_bytes = 0
|
||||
if stream is not None:
|
||||
await stream.close()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Release the model stream owned by this browser session."""
|
||||
await self._close_native_stream()
|
||||
|
||||
async def _emit_transcription(self, text: str, sentence_type: int, end_ms: float, commit_reason: str | None = None) -> None:
|
||||
"""写入或更新一条句子,确保中间结果和最终结果不会在前端产生重复行。"""
|
||||
if not text:
|
||||
return
|
||||
sentence = self.assembler.apply_sentence(
|
||||
{
|
||||
"sentence_id": self.segment_id,
|
||||
"sentence": text,
|
||||
"sentence_type": sentence_type,
|
||||
"start_time": self.segment_start_ms,
|
||||
"end_time": end_ms,
|
||||
"speaker_id": -1,
|
||||
"speaker_name": "",
|
||||
"speaker_evidence": "pending",
|
||||
"speaker_confidence": 0.0,
|
||||
"speaker_strategy": "vllm_no_speaker_evidence",
|
||||
"commit_reason": commit_reason,
|
||||
"speaker_status": ("queued" if sentence_type else "waiting_final") if self.speaker_enabled else "disabled",
|
||||
"speaker_reason": ("等待声纹处理" if sentence_type else "语音片段结束后识别说话人") if self.speaker_enabled else "说话人分离已关闭",
|
||||
}
|
||||
)
|
||||
if sentence_type == 0:
|
||||
self.metrics.partial_count += 1
|
||||
if sentence["revision_count"] > 0:
|
||||
self.metrics.partial_revisions += 1
|
||||
if self.metrics.first_partial_ms is None:
|
||||
self.metrics.first_partial_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
|
||||
else:
|
||||
self.metrics.final_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
|
||||
await self.emit_state(sentence)
|
||||
|
||||
@staticmethod
|
||||
def _is_voice_frame(frame: bytes) -> bool:
|
||||
"""用 PCM 帧的 RMS 判断是否有语音,作为实时低延迟切句触发器。"""
|
||||
if not frame:
|
||||
return False
|
||||
samples = memoryview(frame).cast("h")
|
||||
if not samples:
|
||||
return False
|
||||
square_mean = sum(sample * sample for sample in samples) / len(samples)
|
||||
return math.sqrt(square_mean) >= VAD_RMS_THRESHOLD
|
||||
|
||||
def _strip_wav_header(self, chunk: bytes) -> bytes:
|
||||
"""增量解析 RIFF chunk;支持扩展头,并拒绝采样率或声道不匹配的 WAV。"""
|
||||
if self.wav_payload_started:
|
||||
if self.wav_data_remaining is None:
|
||||
return chunk
|
||||
payload = chunk[:self.wav_data_remaining]
|
||||
self.wav_data_remaining -= len(payload)
|
||||
return payload
|
||||
self.wav_header_buffer.extend(chunk)
|
||||
buffer = self.wav_header_buffer
|
||||
if not self.wav_riff_read:
|
||||
if len(buffer) < 12:
|
||||
return b""
|
||||
if buffer[:4] != b"RIFF" or buffer[8:12] != b"WAVE":
|
||||
raise ValueError("文件不是有效的 RIFF/WAV 音频")
|
||||
del buffer[:12]
|
||||
self.wav_riff_read = True
|
||||
while len(buffer) >= 8:
|
||||
kind = bytes(buffer[:4])
|
||||
size = int.from_bytes(buffer[4:8], "little")
|
||||
if kind == b"data":
|
||||
if not self.wav_format_valid or size % 2:
|
||||
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
|
||||
self.wav_data_remaining = size
|
||||
self.wav_payload_started = True
|
||||
payload = bytes(buffer[8:8 + size])
|
||||
self.wav_data_remaining -= len(payload)
|
||||
buffer.clear()
|
||||
return payload
|
||||
if size > 1024 * 1024:
|
||||
raise ValueError("WAV 元数据头过大,请转换为标准 PCM WAV")
|
||||
chunk_size = 8 + size + (size % 2)
|
||||
if len(buffer) < chunk_size:
|
||||
return b""
|
||||
if kind == b"fmt ":
|
||||
fmt = buffer[8:8 + size]
|
||||
fields = (int.from_bytes(fmt[0:2], "little"), int.from_bytes(fmt[2:4], "little"),
|
||||
int.from_bytes(fmt[4:8], "little"), int.from_bytes(fmt[14:16], "little"))
|
||||
if size < 16 or fields != (1, 1, 16000, 16):
|
||||
raise ValueError("WAV 必须为 16kHz、单声道、PCM16,请先转换音频")
|
||||
self.wav_format_valid = True
|
||||
del buffer[:chunk_size]
|
||||
return b""
|
||||
|
||||
async def _resolve_speaker(self, job: SpeakerJob) -> None:
|
||||
"""异步解析单个 turn 的说话人,并把结果覆盖回同一个 sentence_id。"""
|
||||
if not self.speaker_enabled:
|
||||
return
|
||||
async def update_status(status: str, reason: str) -> None:
|
||||
"""将每个失败或等待阶段回写原片段,避免只发一次全局告警。"""
|
||||
updated = self.assembler.apply_speaker_update({
|
||||
"sentence_id": job.sentence_id, "speaker_id": -1,
|
||||
"speaker_evidence": "pending", "speaker_confidence": 0.0,
|
||||
"speaker_status": status, "speaker_reason": reason,
|
||||
})
|
||||
await self.emit_state(updated)
|
||||
|
||||
# 按有效有声帧检查长度,不能让句尾 800ms 静音把短插话伪装成长样本。
|
||||
if job.voiced_ms < MIN_SPEAKER_VOICE_MS:
|
||||
await update_status("insufficient_audio", f"有效语音不足 {MIN_SPEAKER_VOICE_MS}ms,不继承上一位说话人")
|
||||
return
|
||||
if self.auxiliary_service is None:
|
||||
await update_status("service_unavailable", "未配置说话人辅助模型服务")
|
||||
await self.warn_speaker("未配置辅助模型服务,无法执行实时说话人分离")
|
||||
return
|
||||
await update_status("processing", "正在提取声纹并匹配说话人")
|
||||
try:
|
||||
speaker = await self.auxiliary_service.resolve_speaker(
|
||||
job.audio,
|
||||
self.session_id,
|
||||
job.start_time_ms,
|
||||
job.end_time_ms,
|
||||
)
|
||||
except Exception as exc:
|
||||
# 辅助服务异常不能阻断 ASR;当前片段继续保持 pending,方便定位服务问题。
|
||||
LOGGER.exception("speaker resolve failed: session=%s sentence=%s", self.session_id, job.sentence_id)
|
||||
await update_status("service_error", str(exc))
|
||||
await self.warn_speaker(str(exc))
|
||||
return
|
||||
if not speaker:
|
||||
await update_status("no_embedding", "辅助服务未返回可用声纹结果")
|
||||
return
|
||||
update = dict(speaker)
|
||||
update["sentence_id"] = job.sentence_id
|
||||
update["speaker_name"] = str(update.get("speaker_name") or "")
|
||||
updated = self.assembler.apply_speaker_update(update)
|
||||
if updated is not None:
|
||||
LOGGER.info("speaker result: session=%s sentence=%s status=%s strategy=%s", self.session_id,
|
||||
job.sentence_id, updated.get("speaker_status"), updated.get("speaker_strategy"))
|
||||
await self.emit_state(updated)
|
||||
|
||||
async def process_speakers(self) -> None:
|
||||
"""按 turn 顺序串行访问辅助模型,保证在线聚类中心不会乱序更新。"""
|
||||
while True:
|
||||
item = await self.speaker_queue.get()
|
||||
if isinstance(item, EndOfStream):
|
||||
return
|
||||
await self._resolve_speaker(item)
|
||||
|
||||
async def _commit_segment(self, reason: str = "final") -> None:
|
||||
"""在 VAD 检测到一句结束后提交 final,并异步排队当前 turn 的说话人解析。"""
|
||||
if not self.segment_audio or not self.in_speech:
|
||||
return
|
||||
# 去掉句尾触发切段的静音,ASR 与声纹都使用当前片段的真实有效范围。
|
||||
trailing_bytes = int(self.silence_ms * PARTIAL_BYTES_PER_SECOND / 1000)
|
||||
if trailing_bytes:
|
||||
del self.segment_audio[-trailing_bytes:]
|
||||
final_audio = bytes(self.segment_audio)
|
||||
final_start_ms = self.segment_start_ms
|
||||
final_end_ms = self.segment_start_ms + self._duration_ms()
|
||||
final_sentence_id = self.segment_id
|
||||
text = await self._finish_transcription()
|
||||
if text:
|
||||
await self._emit_transcription(text, 1, final_end_ms, reason)
|
||||
if self.speaker_enabled:
|
||||
await self.speaker_queue.put(
|
||||
SpeakerJob(
|
||||
sentence_id=final_sentence_id,
|
||||
audio=final_audio,
|
||||
start_time_ms=final_start_ms,
|
||||
end_time_ms=final_end_ms,
|
||||
voiced_ms=self.voiced_ms,
|
||||
)
|
||||
)
|
||||
elif self.assembler.segments.pop(final_sentence_id, None) is not None:
|
||||
# final 判定无文本时撤回临时结果,不能留下永远等待声纹的 partial。
|
||||
await self.emit_state()
|
||||
self.segment_audio.clear()
|
||||
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
self.segment_id += 1
|
||||
self.silence_ms = 0
|
||||
self.in_speech = False
|
||||
self.voiced_ms = 0.0
|
||||
|
||||
async def process_audio(self) -> None:
|
||||
"""消费音频,以 VAD 静音结束作为切句主逻辑,并按窗口发送 partial。"""
|
||||
last_partial_bytes = 0
|
||||
partial_bytes = int(self.partial_interval_ms / 1000 * PARTIAL_BYTES_PER_SECOND)
|
||||
while True:
|
||||
item = await self.audio_queue.get()
|
||||
if isinstance(item, EndOfStream):
|
||||
break
|
||||
chunk = self._strip_wav_header(item) if self.source == "file" else item
|
||||
self.metrics.input_chunks += 1
|
||||
if not self.windowed_partial:
|
||||
# websocket_handler 已拒绝压缩文件;这里保留防御分支,避免未来
|
||||
# 新客户端绕过入口时又悄悄退化成“整段上传后切片”。
|
||||
raise RuntimeError("实时流式模式只接受 16kHz PCM16 音频")
|
||||
if not chunk:
|
||||
continue
|
||||
self.vad_buffer.extend(chunk)
|
||||
while len(self.vad_buffer) >= VAD_FRAME_BYTES:
|
||||
frame = bytes(self.vad_buffer[:VAD_FRAME_BYTES])
|
||||
del self.vad_buffer[:VAD_FRAME_BYTES]
|
||||
frame_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
self.processed_audio_bytes += len(frame)
|
||||
self.metrics.audio_bytes += len(frame)
|
||||
voiced = self._is_voice_frame(frame)
|
||||
if voiced and not self.in_speech:
|
||||
self.in_speech = True
|
||||
self.segment_start_ms = frame_start_ms - len(self.pre_roll) / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
self.segment_audio = bytearray(self.pre_roll)
|
||||
self.pre_roll.clear()
|
||||
last_partial_bytes = 0
|
||||
await self._start_native_stream()
|
||||
await self._feed_native_audio()
|
||||
if self.in_speech:
|
||||
self.segment_audio.extend(frame)
|
||||
await self._feed_native_audio()
|
||||
if voiced:
|
||||
self.voiced_ms += VAD_FRAME_MS
|
||||
self.silence_ms = 0 if voiced else self.silence_ms + VAD_FRAME_MS
|
||||
if len(self.segment_audio) - last_partial_bytes >= partial_bytes and self.silence_ms < self.silence_limit_ms:
|
||||
text = await self._partial_text()
|
||||
if text:
|
||||
await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms())
|
||||
last_partial_bytes = len(self.segment_audio)
|
||||
# 短交接停顿优先于普通 800/1400ms 静音切段,但必须先有
|
||||
# 至少 800ms 有效语音,避免把咳嗽、噪声或极短插话送去聚类。
|
||||
short_speaker_gap = (
|
||||
self.speaker_gap_enabled
|
||||
and self.voiced_ms >= MIN_SPEAKER_VOICE_MS
|
||||
and self.silence_ms >= self.speaker_gap_ms
|
||||
)
|
||||
if short_speaker_gap or self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000:
|
||||
reason = (
|
||||
"speaker_gap" if short_speaker_gap
|
||||
else "silence" if self.silence_ms >= self.silence_limit_ms
|
||||
else "max_duration"
|
||||
)
|
||||
await self._commit_segment(reason)
|
||||
last_partial_bytes = 0
|
||||
else:
|
||||
# 参考原 WebSocket 保留 200ms 前滚,减少首字低能量音素被裁掉。
|
||||
self.pre_roll.extend(frame)
|
||||
del self.pre_roll[:-6400]
|
||||
|
||||
if not self.wav_payload_started or (not self.input_stopped and self.wav_data_remaining not in (None, 0)):
|
||||
raise ValueError("WAV 文件不完整,未收到全部音频数据")
|
||||
if len(self.vad_buffer) % 2:
|
||||
raise ValueError("PCM16 音频必须包含完整的双字节采样")
|
||||
if self.windowed_partial and self.vad_buffer:
|
||||
tail_ms = len(self.vad_buffer) / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
tail_voiced = self._is_voice_frame(bytes(self.vad_buffer))
|
||||
if tail_voiced and not self.in_speech:
|
||||
self.in_speech = True
|
||||
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
|
||||
self.processed_audio_bytes += len(self.vad_buffer)
|
||||
self.metrics.audio_bytes += len(self.vad_buffer)
|
||||
if self.in_speech:
|
||||
self.segment_audio.extend(self.vad_buffer)
|
||||
self.silence_ms = 0 if tail_voiced else self.silence_ms + tail_ms
|
||||
self.voiced_ms += tail_ms if tail_voiced else 0
|
||||
self.vad_buffer.clear()
|
||||
if self.segment_audio:
|
||||
if not self.windowed_partial:
|
||||
self.in_speech = True
|
||||
await self._commit_segment()
|
||||
await self.emit({"type": "metrics", "metrics": self.metrics.snapshot()})
|
||||
|
||||
|
||||
def deployment_model_name() -> str:
|
||||
"""把部署脚本的 0.6b/1.7b 别名解析成 vLLM 对外发布的模型名。"""
|
||||
public_name = os.getenv("VLLM_SERVED_MODEL_NAME")
|
||||
if public_name:
|
||||
return public_name
|
||||
requested = os.getenv("QWEN3_ASR_MODEL", "default")
|
||||
with (DEPLOY_ROOT / "model_manifest.json").open(encoding="utf-8") as source:
|
||||
manifest = json.load(source)
|
||||
if requested == "default":
|
||||
return str(manifest["default_model"])
|
||||
for model_id, config in manifest["models"].items():
|
||||
if requested.lower() == str(config.get("alias", "")).lower():
|
||||
return model_id
|
||||
return requested
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
"""只解析服务选择参数;浏览器服务端口继续由代码内部变量统一维护。"""
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-service-url", default=os.getenv("MODEL_SERVICE_URL", DEFAULT_MODEL_SERVICE_URL))
|
||||
parser.add_argument("--model", default=deployment_model_name())
|
||||
parser.add_argument("--no-browser", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
async def index_handler(_: web.Request) -> web.FileResponse:
|
||||
"""返回独立 Demo 测试页面,并避免入口页缓存旧的脚本版本号。"""
|
||||
# 入口页必须每次重新校验,配合 app.js 的版本号变更,避免用户继续运行旧前端。
|
||||
return web.FileResponse(
|
||||
Path(__file__).parent / "static" / "index.html",
|
||||
headers={"Cache-Control": "no-store"},
|
||||
)
|
||||
|
||||
|
||||
async def config_handler(request: web.Request) -> web.Response:
|
||||
"""暴露服务启动时的默认配置,让页面自动填充 VLLM 地址和模型名。"""
|
||||
config = request.app[MODEL_SERVICE_KEY].config
|
||||
auxiliary = request.app.get(AUXILIARY_SERVICE_KEY)
|
||||
return web.json_response({
|
||||
"model_service_url": config.base_url, "model": config.model,
|
||||
"speaker_service_url": getattr(getattr(auxiliary, "config", None), "base_url", None),
|
||||
})
|
||||
|
||||
|
||||
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
||||
"""处理一个浏览器会话,每个连接独立保存音频、句子和展示状态。"""
|
||||
ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024)
|
||||
await ws.prepare(request)
|
||||
default_model_service: VLLMTranscriptionService = request.app[MODEL_SERVICE_KEY]
|
||||
model_service = default_model_service
|
||||
owns_model_service = False
|
||||
processing: asyncio.Task[None] | None = None
|
||||
speaker_processing: asyncio.Task[None] | None = None
|
||||
session: RealtimeSession | None = None
|
||||
try:
|
||||
first = await ws.receive()
|
||||
if first.type != WSMsgType.TEXT:
|
||||
await ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||||
return ws
|
||||
try:
|
||||
start = json.loads(first.data)
|
||||
except json.JSONDecodeError:
|
||||
await ws.send_json({"type": "error", "message": "invalid start JSON"})
|
||||
return ws
|
||||
if not isinstance(start, dict) or start.get("type") != "start":
|
||||
await ws.send_json({"type": "error", "message": "first message must have type=start"})
|
||||
return ws
|
||||
|
||||
# 本项目用于验证实时流式链路,文件模式只接受可以按 PCM 帧连续处理的
|
||||
# WAV/PCM;MP3、M4A 等压缩容器只能在文件完整到达后解码,不纳入本次测试。
|
||||
source = str(start.get("source") or "mic")
|
||||
file_suffix = Path(str(start.get("file_name") or "")).suffix.lower()
|
||||
if source == "file" and file_suffix not in {".pcm", ".wav"}:
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "error",
|
||||
"message": "实时流式测试的文件模式只支持 PCM 或 WAV,请改用麦克风、PCM 或 WAV",
|
||||
}
|
||||
)
|
||||
return ws
|
||||
|
||||
# 页面可以在不重启 WebSocket Demo 的情况下为当前会话切换 VLLM 地址;
|
||||
# 未切换时继续复用默认服务,避免普通场景为每个连接重复创建 HTTP 会话。
|
||||
try:
|
||||
requested_url = validate_model_service_url(
|
||||
str(start.get("model_service_url") or default_model_service.config.base_url)
|
||||
)
|
||||
except ValueError as exc:
|
||||
await ws.send_json({"type": "error", "message": str(exc)})
|
||||
return ws
|
||||
requested_model = str(start.get("model") or default_model_service.config.model).strip()
|
||||
if (
|
||||
requested_url != default_model_service.config.base_url
|
||||
or requested_model != default_model_service.config.model
|
||||
):
|
||||
model_service = VLLMTranscriptionService(
|
||||
ModelServiceConfig(base_url=requested_url, model=requested_model)
|
||||
)
|
||||
await model_service.start()
|
||||
owns_model_service = True
|
||||
|
||||
auxiliary_service = request.app.get(AUXILIARY_SERVICE_KEY)
|
||||
session = RealtimeSession(ws, model_service, auxiliary_service, start)
|
||||
auxiliary_config = getattr(auxiliary_service, "config", None)
|
||||
speaker_health: dict[str, Any] | None = None
|
||||
speaker_health_error: str | None = None
|
||||
if session.speaker_enabled and auxiliary_service is not None:
|
||||
# 健康检查只用于尽早暴露辅助服务问题;即使失败也不阻断 ASR,
|
||||
# 这样可以从同一页面继续观察 ASR 与说话人链路的差异。
|
||||
health_check = getattr(auxiliary_service, "health", None)
|
||||
if callable(health_check):
|
||||
try:
|
||||
speaker_health = await asyncio.wait_for(health_check(), timeout=5)
|
||||
if speaker_health.get("speaker_embedding_ready", speaker_health.get("ready")) is False:
|
||||
speaker_health_error = "辅助模型服务未就绪,请检查 /health 返回的 models 状态"
|
||||
except Exception as exc:
|
||||
speaker_health_error = f"说话人辅助服务不可用:{exc}"
|
||||
if speaker_health_error:
|
||||
# 辅助服务未就绪时先保留原始 VAD 切段,避免在没有声纹结果的
|
||||
# 情况下增加大量短片段;服务恢复后由新的会话重新启用。
|
||||
session.speaker_gap_enabled = False
|
||||
await session.emit(
|
||||
{
|
||||
"type": "start",
|
||||
"model_service_url": model_service.config.base_url,
|
||||
"model": model_service.config.model,
|
||||
"session_id": session.session_id,
|
||||
"enable_native_partial_stream": session.enable_native_partial,
|
||||
"native_partial_supported": model_service.native_partial_supported,
|
||||
"partial_mode": "vllm_realtime_websocket" if session.enable_native_partial and model_service.native_partial_supported else "http_cumulative_window",
|
||||
"speaker_diarization_enabled": session.speaker_enabled,
|
||||
"speaker_service_url": getattr(auxiliary_config, "base_url", None),
|
||||
"speaker_service_health": speaker_health,
|
||||
"speaker_gap_enabled": session.speaker_gap_enabled,
|
||||
"speaker_gap_ms": session.speaker_gap_ms,
|
||||
"sentence_strategy": session.sentence_strategy,
|
||||
"silence_limit_ms": session.silence_limit_ms,
|
||||
"display_state_supported": True,
|
||||
}
|
||||
)
|
||||
if speaker_health_error:
|
||||
await session.warn_speaker(speaker_health_error)
|
||||
processing = asyncio.create_task(session.process_audio())
|
||||
if session.speaker_enabled:
|
||||
# speaker worker 与音频处理并行运行;它只消费已经结束的 turn,
|
||||
# 因此不会阻塞下一帧音频进入队列或影响 ASR partial 输出。
|
||||
speaker_processing = asyncio.create_task(session.process_speakers())
|
||||
async def guarded(operation):
|
||||
"""接收和队列背压同时监听 worker,推理失败立即报错而非永远等 stop。"""
|
||||
pending = asyncio.create_task(operation)
|
||||
try:
|
||||
workers = [task for task in (processing, speaker_processing) if task is not None]
|
||||
done, _ = await asyncio.wait([pending, *workers], return_when=asyncio.FIRST_COMPLETED)
|
||||
if pending in done:
|
||||
return await pending
|
||||
for worker in workers:
|
||||
if worker in done:
|
||||
await worker
|
||||
raise RuntimeError("实时处理任务意外结束")
|
||||
return await pending
|
||||
finally:
|
||||
if not pending.done():
|
||||
pending.cancel()
|
||||
await asyncio.gather(pending, return_exceptions=True)
|
||||
|
||||
input_finished = False
|
||||
while not ws.closed:
|
||||
message = await guarded(ws.receive())
|
||||
if message.type == WSMsgType.BINARY:
|
||||
await guarded(session.audio_queue.put(bytes(message.data)))
|
||||
continue
|
||||
if message.type == WSMsgType.TEXT:
|
||||
try:
|
||||
control = json.loads(message.data)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not isinstance(control, dict):
|
||||
continue
|
||||
if control.get("type") in {"eof", "stop"}:
|
||||
input_finished = True
|
||||
session.input_stopped = control.get("type") == "stop"
|
||||
await session.emit({"type": "draining", "message": "正在完成转写和说话人识别"})
|
||||
await guarded(session.audio_queue.put(EOF))
|
||||
break
|
||||
if control.get("type") == "abort":
|
||||
processing.cancel()
|
||||
if speaker_processing is not None:
|
||||
speaker_processing.cancel()
|
||||
await asyncio.gather(
|
||||
processing,
|
||||
*(task for task in [speaker_processing] if task is not None),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return ws
|
||||
if message.type in {WSMsgType.ERROR, WSMsgType.CLOSE, WSMsgType.CLOSED}:
|
||||
processing.cancel()
|
||||
if speaker_processing is not None:
|
||||
speaker_processing.cancel()
|
||||
await asyncio.gather(
|
||||
processing,
|
||||
*(task for task in [speaker_processing] if task is not None),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return ws
|
||||
if not input_finished:
|
||||
# 浏览器或网络断开后,音频生产者已经不存在,不能让处理任务继续等待
|
||||
# 永远不会到来的 EOF,因此这里主动取消任务并回收异常结果。
|
||||
processing.cancel()
|
||||
if speaker_processing is not None:
|
||||
speaker_processing.cancel()
|
||||
await asyncio.gather(
|
||||
processing,
|
||||
*(task for task in [speaker_processing] if task is not None),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return ws
|
||||
try:
|
||||
await processing
|
||||
if speaker_processing is not None:
|
||||
await session.speaker_queue.put(EOF)
|
||||
await speaker_processing
|
||||
await session.emit_state()
|
||||
await session.emit({
|
||||
"type": "end", "metrics": session.metrics.snapshot(),
|
||||
"sentences": session.assembler.raw_snapshot(),
|
||||
"display_blocks": session.assembler.display_blocks(session.merge_adjacent),
|
||||
})
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
await session.emit({"type": "error", "message": str(exc)})
|
||||
except Exception as exc:
|
||||
LOGGER.exception("WebSocket session failed")
|
||||
if not ws.closed:
|
||||
await ws.send_json({"type": "error", "message": str(exc)})
|
||||
finally:
|
||||
for task in (processing, speaker_processing):
|
||||
if task is not None and not task.done():
|
||||
task.cancel()
|
||||
pending_tasks = [task for task in (processing, speaker_processing) if task is not None]
|
||||
if pending_tasks:
|
||||
await asyncio.gather(*pending_tasks, return_exceptions=True)
|
||||
# stop、abort、断线和推理异常均释放会话,清理失败不覆盖最终识别结果。
|
||||
if session is not None:
|
||||
reset = getattr(session.auxiliary_service, "reset_speaker_session", None)
|
||||
await session.close()
|
||||
if reset is not None:
|
||||
try:
|
||||
await asyncio.wait_for(reset(session.session_id), timeout=5)
|
||||
except Exception:
|
||||
LOGGER.warning("speaker session cleanup failed: %s", session.session_id, exc_info=True)
|
||||
if owns_model_service:
|
||||
await model_service.close()
|
||||
if not ws.closed:
|
||||
await ws.close()
|
||||
return ws
|
||||
|
||||
|
||||
async def start_app(model_service_url: str, model: str) -> web.Application:
|
||||
"""创建 HTTP/WebSocket 应用,并挂载可复用的 VLLM 适配器。"""
|
||||
app = web.Application()
|
||||
app[MODEL_SERVICE_KEY] = VLLMTranscriptionService(
|
||||
ModelServiceConfig(
|
||||
base_url=model_service_url,
|
||||
model=model,
|
||||
realtime_enabled=os.getenv("VLLM_REALTIME", "true").lower() not in {"0", "false", "no", "off"},
|
||||
)
|
||||
)
|
||||
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
|
||||
AuxiliaryServiceConfig(base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"))
|
||||
)
|
||||
|
||||
async def lifecycle(application: web.Application):
|
||||
await application[MODEL_SERVICE_KEY].start()
|
||||
await application[AUXILIARY_SERVICE_KEY].start()
|
||||
yield
|
||||
await application[AUXILIARY_SERVICE_KEY].close()
|
||||
await application[MODEL_SERVICE_KEY].close()
|
||||
|
||||
app.cleanup_ctx.append(lifecycle)
|
||||
app.router.add_get("/", index_handler)
|
||||
app.router.add_get("/api/config", config_handler)
|
||||
app.router.add_static("/static/", Path(__file__).parent / "static")
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
# 参考腾讯 Demo 将 static 目录挂载到根路径;页面中的 style.css 和 app.js
|
||||
# 使用相对地址,必须同时提供根路径静态资源路由,否则浏览器会显示无样式页面。
|
||||
app.router.add_static("/", Path(__file__).parent / "static", show_index=False)
|
||||
return app
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""启动本地测试页面和 WebSocket 服务。"""
|
||||
args = parse_args()
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
if not args.no_browser:
|
||||
webbrowser.open(f"http://{WEB_DISPLAY_HOST}:{WEB_PORT}/")
|
||||
print(f"WebSocket demo: http://{WEB_DISPLAY_HOST}:{WEB_PORT}/", flush=True)
|
||||
print(f"VLLM service: {args.model_service_url} ({args.model})", flush=True)
|
||||
web.run_app(start_app(args.model_service_url, args.model), host=WEB_HOST, port=WEB_PORT)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,101 +0,0 @@
|
|||
"""独立音频请求准备逻辑的回归测试。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
import unittest
|
||||
import wave
|
||||
from io import BytesIO
|
||||
|
||||
from model_service import (
|
||||
pcm16_to_wav,
|
||||
prepare_audio_request,
|
||||
realtime_ws_url,
|
||||
wav_to_pcm16,
|
||||
VLLMRealtimeStream,
|
||||
)
|
||||
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||||
from aiohttp import WSMsgType
|
||||
|
||||
|
||||
|
||||
class FakeRealtimeWebSocket:
|
||||
def __init__(self) -> None:
|
||||
self.messages: asyncio.Queue[object] = asyncio.Queue()
|
||||
self.sent: list[dict[str, object]] = []
|
||||
self.closed = False
|
||||
|
||||
def __aiter__(self):
|
||||
return self._messages()
|
||||
|
||||
async def _messages(self):
|
||||
while not self.closed:
|
||||
yield await self.messages.get()
|
||||
|
||||
async def send_json(self, payload: dict[str, object]) -> None:
|
||||
self.sent.append(payload)
|
||||
if payload.get("type") == "input_audio_buffer.commit" and payload.get("final"):
|
||||
await self.messages.put(
|
||||
SimpleNamespace(
|
||||
type=WSMsgType.TEXT,
|
||||
data=json.dumps({"type": "transcription.done", "text": "final text"}),
|
||||
)
|
||||
)
|
||||
|
||||
def exception(self):
|
||||
return None
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
class ModelServiceTests(unittest.TestCase):
|
||||
def test_pcm_is_wrapped_as_16k_mono_wav(self) -> None:
|
||||
wav_bytes = pcm16_to_wav(b"\x00\x00" * 160)
|
||||
with wave.open(BytesIO(wav_bytes), "rb") as wav_file:
|
||||
self.assertEqual(wav_file.getframerate(), 16000)
|
||||
self.assertEqual(wav_file.getnchannels(), 1)
|
||||
self.assertEqual(wav_to_pcm16(wav_bytes), b"\x00\x00" * 160)
|
||||
|
||||
def test_compressed_partial_is_deferred(self) -> None:
|
||||
self.assertIsNone(prepare_audio_request(b"partial", "file", "sample.mp3", partial=True))
|
||||
prepared = prepare_audio_request(b"complete", "file", "sample.mp3", partial=False)
|
||||
self.assertEqual(prepared[1], "sample.mp3")
|
||||
|
||||
def test_realtime_url_uses_the_openai_v1_path(self) -> None:
|
||||
self.assertEqual(
|
||||
realtime_ws_url("https://asr.example/v1/"),
|
||||
"wss://asr.example/v1/realtime",
|
||||
)
|
||||
|
||||
def test_realtime_stream_sends_incremental_audio_and_finishes(self) -> None:
|
||||
async def exercise() -> None:
|
||||
websocket = FakeRealtimeWebSocket()
|
||||
stream = VLLMRealtimeStream(websocket, timeout_seconds=1)
|
||||
await websocket.messages.put(
|
||||
SimpleNamespace(
|
||||
type=WSMsgType.TEXT,
|
||||
data=json.dumps({"type": "transcription.delta", "delta": "partial"}),
|
||||
)
|
||||
)
|
||||
await stream.append_audio(b"\x01\x02")
|
||||
final_text = await stream.finish()
|
||||
self.assertEqual(final_text, "final text")
|
||||
self.assertEqual(
|
||||
base64.b64decode(str(websocket.sent[0]["audio"])),
|
||||
b"\x01\x02",
|
||||
)
|
||||
self.assertEqual(websocket.sent[-1]["type"], "input_audio_buffer.commit")
|
||||
await stream.close()
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
def test_auxiliary_config_is_independent_from_vllm(self) -> None:
|
||||
service = AuxiliaryModelService(AuxiliaryServiceConfig())
|
||||
self.assertEqual(service.config.base_url, "http://127.0.0.1:8010")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -1,388 +0,0 @@
|
|||
"""使用模拟 VLLM 适配器验证本地 WebSocket 流程。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
import asyncio
|
||||
import io
|
||||
import wave
|
||||
import os
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import AioHTTPTestCase
|
||||
|
||||
from server import (
|
||||
AUXILIARY_SERVICE_KEY,
|
||||
MODEL_SERVICE_KEY,
|
||||
config_handler,
|
||||
deployment_model_name,
|
||||
validate_model_service_url,
|
||||
websocket_handler,
|
||||
)
|
||||
|
||||
|
||||
class FakeModelService:
|
||||
native_partial_supported = False
|
||||
config = SimpleNamespace(base_url="http://fake/v1", model="fake-model")
|
||||
|
||||
async def transcribe(self, audio_bytes: bytes, source: str, file_name: str, partial: bool) -> str:
|
||||
return "partial text" if partial else "final text"
|
||||
|
||||
|
||||
class FakeAuxiliaryService:
|
||||
"""返回固定时间段的聚类服务,用于验证 final 后的同句 speaker 更新。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def health(self) -> dict[str, object]:
|
||||
"""模拟辅助模型服务已完成预加载。"""
|
||||
return {"ready": True, "speaker_embedding_ready": True}
|
||||
|
||||
async def resolve_speaker(
|
||||
self,
|
||||
audio_bytes: bytes,
|
||||
session_id: str,
|
||||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> dict[str, object]:
|
||||
_ = (audio_bytes, session_id, start_time_ms, end_time_ms)
|
||||
speaker_id = [0, 1, 0][min(self.calls, 2)]
|
||||
self.calls += 1
|
||||
return {
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_name": f"说话人 {speaker_id + 1}",
|
||||
"speaker_evidence": "fresh",
|
||||
"speaker_confidence": 0.9,
|
||||
"speaker_strategy": "online_embedding_cluster",
|
||||
}
|
||||
|
||||
async def reset_speaker_session(self, session_id: str) -> None:
|
||||
_ = session_id
|
||||
|
||||
|
||||
class UnhealthyAuxiliaryService(FakeAuxiliaryService):
|
||||
"""模拟端口可访问但辅助模型尚未就绪的服务。"""
|
||||
|
||||
async def health(self) -> dict[str, object]:
|
||||
return {"ready": False, "speaker_embedding_ready": False}
|
||||
|
||||
|
||||
class WebSocketFlowTests(AioHTTPTestCase):
|
||||
async def collect(self, ws, until="end"):
|
||||
"""限定等待时间,回归测试中的队列卡死必须表现为失败。"""
|
||||
events = []
|
||||
async with asyncio.timeout(10):
|
||||
while True:
|
||||
event = await ws.receive_json()
|
||||
events.append(event)
|
||||
if event["type"] == until:
|
||||
return events
|
||||
|
||||
def get_app(self) -> web.Application:
|
||||
app = web.Application()
|
||||
app[MODEL_SERVICE_KEY] = FakeModelService()
|
||||
app[AUXILIARY_SERVICE_KEY] = FakeAuxiliaryService()
|
||||
app.router.add_get("/api/config", config_handler)
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
return app
|
||||
|
||||
def test_vllm_url_validation(self) -> None:
|
||||
self.assertEqual(validate_model_service_url(" http://asr.local/v1/ "), "http://asr.local/v1")
|
||||
with self.assertRaises(ValueError):
|
||||
validate_model_service_url("asr.local:8000/v1")
|
||||
with self.assertRaises(ValueError):
|
||||
validate_model_service_url("http://user:password@asr.local/v1")
|
||||
|
||||
def test_deployment_alias_and_served_name_match_vllm(self):
|
||||
"""WebSocket 不能把下载别名直接当成 vLLM 公开模型名。"""
|
||||
with patch.dict(os.environ, {"QWEN3_ASR_MODEL": "0.6b"}, clear=True):
|
||||
self.assertEqual(deployment_model_name(), "Qwen/Qwen3-ASR-0.6B")
|
||||
with patch.dict(os.environ, {"QWEN3_ASR_MODEL": "0.6b", "VLLM_SERVED_MODEL_NAME": "custom-asr"}, clear=True):
|
||||
self.assertEqual(deployment_model_name(), "custom-asr")
|
||||
|
||||
async def test_empty_final_retracts_partial_instead_of_leaving_pending(self):
|
||||
"""最终没有识别文本时撤回临时内容,不留下永远等待声纹的行。"""
|
||||
async def transcribe(*args, partial):
|
||||
return "temporary" if partial else ""
|
||||
self.app[MODEL_SERVICE_KEY].transcribe = transcribe
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "partial_interval_ms": 300})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
await ws.send_json({"type": "eof"})
|
||||
events = await self.collect(ws)
|
||||
self.assertTrue(any(e["type"] == "sentences" for e in events))
|
||||
self.assertEqual(events[-1]["sentences"], [])
|
||||
self.assertEqual(events[-1]["display_blocks"], [])
|
||||
|
||||
async def test_compressed_file_is_rejected_for_streaming_validation(self) -> None:
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "source": "file", "file_name": "meeting.mp3"})
|
||||
error = await ws.receive_json()
|
||||
self.assertEqual(error["type"], "error")
|
||||
self.assertIn("PCM 或 WAV", error["message"])
|
||||
await ws.close()
|
||||
|
||||
async def test_short_interruption_does_not_inherit_or_call_embedding(self):
|
||||
"""句尾静音不能凑够声纹时长,长段 A 后的短插话独立保持 pending。"""
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "source": "mic"})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000 + b"\x00\x00" * 12800 + b"\xe8\x03" * 8000 + b"\x00\x00" * 12800)
|
||||
await ws.send_json({"type": "stop"})
|
||||
events = await self.collect(ws)
|
||||
final = events[-1]
|
||||
self.assertEqual(self.app[AUXILIARY_SERVICE_KEY].calls, 1)
|
||||
self.assertEqual([s["speaker_id"] for s in final["sentences"]], [0, -1])
|
||||
self.assertEqual(final["sentences"][1]["speaker_status"], "insufficient_audio")
|
||||
self.assertEqual(final["sentences"][0]["end_time"], 1000)
|
||||
self.assertEqual(len(final["display_blocks"]), 2)
|
||||
|
||||
async def test_short_speaker_gap_splits_turn_before_next_speaker(self):
|
||||
"""说话人模式在短交接停顿处切段,保持每段独立进入在线聚类。"""
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
"speaker_diarization": 1,
|
||||
"speaker_gap_ms": 400,
|
||||
"partial_interval_ms": 1200,
|
||||
}
|
||||
)
|
||||
self.assertEqual((await ws.receive_json())["type"], "start")
|
||||
|
||||
voiced = b"\xe8\x03" * 16000 # 每位说话人一秒有效语音
|
||||
handoff_gap = b"\x00\x00" * 6400 # 400ms,低于原始 800ms 切段阈值
|
||||
trailing_silence = b"\x00\x00" * 12800
|
||||
await ws.send_bytes(voiced + handoff_gap + voiced + trailing_silence)
|
||||
await ws.send_json({"type": "eof"})
|
||||
|
||||
events: list[dict[str, object]] = []
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
events.append(message)
|
||||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
final_sentences = [
|
||||
message["sentences"][0]
|
||||
for message in events
|
||||
if message["type"] == "sentences" and message["sentences"][0]["sentence_type"] == 1
|
||||
]
|
||||
latest_by_id = {int(item["sentence_id"]): item for item in final_sentences}
|
||||
latest = [latest_by_id[index] for index in sorted(latest_by_id)]
|
||||
self.assertEqual(len(latest), 2)
|
||||
self.assertEqual([item["speaker_id"] for item in latest], [0, 1])
|
||||
self.assertEqual(self.app[AUXILIARY_SERVICE_KEY].calls, 2)
|
||||
await ws.close()
|
||||
|
||||
async def test_stop_waits_for_slow_speaker_and_includes_final_snapshot(self):
|
||||
"""在 end 前必须收到所有声纹结果,不能复现页面原先五秒断开的行为。"""
|
||||
auxiliary = self.app[AUXILIARY_SERVICE_KEY]
|
||||
original = auxiliary.resolve_speaker
|
||||
async def delayed(*args):
|
||||
await asyncio.sleep(5.1)
|
||||
return await original(*args)
|
||||
auxiliary.resolve_speaker = delayed
|
||||
auxiliary.reset_speaker_session = AsyncMock()
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start"})
|
||||
start = await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
await ws.send_json({"type": "stop"})
|
||||
events = await self.collect(ws)
|
||||
self.assertEqual(events[-1]["sentences"][0]["speaker_id"], 0)
|
||||
self.assertTrue(any(e["type"] == "draining" for e in events))
|
||||
self.assertTrue(any(e["type"] == "sentences" and e["sentences"][0]["speaker_status"] == "processing" for e in events))
|
||||
await ws.receive() # 等待服务端执行 finally 并关闭连接
|
||||
auxiliary.reset_speaker_session.assert_awaited_once_with(start["session_id"])
|
||||
|
||||
async def test_speaker_error_is_visible_on_segment_and_asr_finishes(self):
|
||||
"""声纹推理失败不能吞掉转写,且每条失败片段要携带诊断原因。"""
|
||||
self.app[AUXILIARY_SERVICE_KEY].resolve_speaker = AsyncMock(side_effect=RuntimeError("embedding model missing"))
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start"})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
await ws.send_json({"type": "eof"})
|
||||
events = await self.collect(ws)
|
||||
segment = events[-1]["sentences"][0]
|
||||
self.assertEqual(segment["sentence"], "final text")
|
||||
self.assertEqual(segment["speaker_status"], "service_error")
|
||||
self.assertIn("embedding model missing", segment["speaker_reason"])
|
||||
self.assertTrue(any(e["type"] == "speaker_warning" for e in events))
|
||||
|
||||
async def test_asr_failure_is_reported_before_stop(self):
|
||||
"""音频 worker 抛异常时,接收任务应立即报告,不能等到客户端 stop。"""
|
||||
self.app[MODEL_SERVICE_KEY].transcribe = AsyncMock(side_effect=RuntimeError("vllm unavailable"))
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "partial_interval_ms": 300})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
events = await self.collect(ws, until="error")
|
||||
self.assertIn("vllm unavailable", events[-1]["message"])
|
||||
|
||||
async def test_unknown_speaker_response_is_diagnosable(self):
|
||||
"""旧服务只回标签、缺少 fresh/confidence 时应说明拒绝原因。"""
|
||||
self.app[AUXILIARY_SERVICE_KEY].resolve_speaker = AsyncMock(return_value={"speaker_id": 0})
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start"})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
await ws.send_json({"type": "eof"})
|
||||
events = await self.collect(ws)
|
||||
self.assertEqual(events[-1]["sentences"][0]["speaker_status"], "evidence_rejected")
|
||||
|
||||
async def test_abort_and_disconnect_release_cluster_state(self):
|
||||
"""清理不应只存在于成功 stop 的路径。"""
|
||||
auxiliary = self.app[AUXILIARY_SERVICE_KEY]
|
||||
auxiliary.reset_speaker_session = AsyncMock()
|
||||
for abort in (True, False):
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start"})
|
||||
await ws.receive_json()
|
||||
if abort:
|
||||
await ws.send_json({"type": "abort"})
|
||||
await ws.receive()
|
||||
else:
|
||||
await ws.close()
|
||||
async with asyncio.timeout(2):
|
||||
while auxiliary.reset_speaker_session.await_count < 2:
|
||||
await asyncio.sleep(0.01)
|
||||
self.assertEqual(auxiliary.reset_speaker_session.await_count, 2)
|
||||
|
||||
async def test_extended_wav_header_is_removed_before_asr(self):
|
||||
"""分片 RIFF/JUNK/fmt/data 头不能混入声纹和 ASR 的 PCM 数据。"""
|
||||
output = io.BytesIO()
|
||||
pcm = b"\xe8\x03" * 16000
|
||||
with wave.open(output, "wb") as wav:
|
||||
wav.setparams((1, 2, 16000, 0, "NONE", ""))
|
||||
wav.writeframes(pcm)
|
||||
original = output.getvalue()
|
||||
junk = b"JUNK\x04\x00\x00\x00test"
|
||||
payload = b"RIFF" + (len(original) - 8 + len(junk)).to_bytes(4, "little") + original[8:12] + junk + original[12:]
|
||||
transcribe = AsyncMock(return_value="wav text")
|
||||
self.app[MODEL_SERVICE_KEY].transcribe = transcribe
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "source": "file", "file_name": "test.wav"})
|
||||
await ws.receive_json()
|
||||
for offset in range(0, len(payload), 337):
|
||||
await ws.send_bytes(payload[offset:offset + 337])
|
||||
await ws.send_json({"type": "eof"})
|
||||
events = await self.collect(ws)
|
||||
self.assertEqual(transcribe.call_args.args[0], pcm)
|
||||
self.assertEqual(events[-1]["sentences"][0]["end_time"], 1000)
|
||||
|
||||
async def test_incompatible_wav_is_rejected_immediately(self):
|
||||
"""非 16kHz 单声道 WAV 不能被误解释为可识别的 PCM16。"""
|
||||
output = io.BytesIO()
|
||||
with wave.open(output, "wb") as wav:
|
||||
wav.setparams((2, 2, 44100, 0, "NONE", ""))
|
||||
wav.writeframes(b"\x00\x00" * 2000)
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "source": "file", "file_name": "bad.wav"})
|
||||
await ws.receive_json()
|
||||
await ws.send_bytes(output.getvalue())
|
||||
events = await self.collect(ws, until="error")
|
||||
self.assertIn("16kHz", events[-1]["message"])
|
||||
|
||||
async def test_frontend_can_read_default_vllm_config(self) -> None:
|
||||
response = await self.client.get("/api/config")
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(await response.json(), {"model_service_url": "http://fake/v1", "model": "fake-model", "speaker_service_url": None})
|
||||
|
||||
async def test_partial_and_final_share_one_sentence_id(self) -> None:
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
"speaker_diarization": 0,
|
||||
"partial_interval_ms": 300,
|
||||
"max_segment_sec": 12,
|
||||
}
|
||||
)
|
||||
start = await ws.receive_json()
|
||||
self.assertEqual(start["type"], "start")
|
||||
# 使用幅度足够的 PCM 语音帧,静音帧会被新切句器正确忽略。
|
||||
await ws.send_bytes(b"\xe8\x03" * 16000)
|
||||
|
||||
messages = []
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
messages.append(message)
|
||||
if message["type"] == "sentences":
|
||||
self.assertEqual(message["sentences"][0]["sentence_id"], 0)
|
||||
if message["sentences"][0]["sentence_type"] == 0:
|
||||
break
|
||||
|
||||
await ws.send_json({"type": "eof"})
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
messages.append(message)
|
||||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
sentence_events = [message for message in messages if message["type"] == "sentences"]
|
||||
self.assertGreaterEqual(len(sentence_events), 2)
|
||||
self.assertTrue(all(event["sentences"][0]["sentence_id"] == 0 for event in sentence_events))
|
||||
self.assertTrue(all(event["sentences"][0]["sentence_type"] == 0 for event in sentence_events[:-1]))
|
||||
self.assertEqual(sentence_events[-1]["sentences"][0]["sentence_type"], 1)
|
||||
self.assertEqual(sentence_events[-1]["sentences"][0]["sentence"], "final text")
|
||||
await ws.close()
|
||||
|
||||
async def test_speaker_health_failure_is_reported_without_blocking_asr(self) -> None:
|
||||
"""辅助模型未就绪时先报告告警,同时保留 ASR 会话能力。"""
|
||||
self.app[AUXILIARY_SERVICE_KEY] = UnhealthyAuxiliaryService()
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "source": "mic", "speaker_diarization": 1})
|
||||
|
||||
start = await ws.receive_json()
|
||||
warning = await ws.receive_json()
|
||||
|
||||
self.assertEqual(start["type"], "start")
|
||||
self.assertFalse(start["speaker_service_health"]["ready"])
|
||||
self.assertFalse(start["speaker_gap_enabled"])
|
||||
self.assertEqual(warning["type"], "speaker_warning")
|
||||
self.assertIn("未就绪", warning["message"])
|
||||
await ws.close()
|
||||
|
||||
async def test_vad_split_and_speaker_update(self) -> None:
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
"speaker_diarization": 1,
|
||||
"partial_interval_ms": 1200,
|
||||
}
|
||||
)
|
||||
self.assertEqual((await ws.receive_json())["type"], "start")
|
||||
|
||||
voiced = b"\xe8\x03" * 16000 # 每段一秒有效语音,满足独立声纹长度要求
|
||||
silence = b"\x00\x00" * 12800 # 0.8 秒静音,触发当前 turn 提交
|
||||
await ws.send_bytes(voiced + silence + voiced + silence + voiced)
|
||||
await ws.send_json({"type": "eof"})
|
||||
|
||||
events: list[dict[str, object]] = []
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
events.append(message)
|
||||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
sentence_events = [message for message in events if message["type"] == "sentences"]
|
||||
final_sentences = [
|
||||
message["sentences"][0]
|
||||
for message in sentence_events
|
||||
if message["sentences"][0]["sentence_type"] == 1
|
||||
]
|
||||
latest_by_id = {int(item["sentence_id"]): item for item in final_sentences}
|
||||
self.assertGreaterEqual(len(latest_by_id), 2)
|
||||
latest = [latest_by_id[index] for index in sorted(latest_by_id)[-3:]]
|
||||
self.assertEqual([item["speaker_id"] for item in latest], [0, 1, 0])
|
||||
self.assertTrue(all(item["speaker_evidence"] == "fresh" for item in latest))
|
||||
await ws.close()
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
"""Unit tests for the migrated FunASR streaming lifecycle."""
|
||||
"""FunASR 流式生命周期迁移测试。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
"""Contract test for the unchanged Tencent UI to FunASR native WS bridge."""
|
||||
"""验证未修改的腾讯界面与 FunASR 原生 WebSocket 桥接协议。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -18,7 +18,7 @@ from backend.realtime_websocket.funasr_server import (
|
|||
|
||||
|
||||
class FakeNativeWebSocket:
|
||||
"""Stand in for FunASR's native WSS process without loading model weights."""
|
||||
"""用模拟进程代替 FunASR 原生 WSS 服务,不加载模型权重。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.incoming: asyncio.Queue[str] = asyncio.Queue()
|
||||
|
|
@ -86,7 +86,7 @@ class FunASRBridgeTests(AioHTTPTestCase):
|
|||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
# The server keeps speaker labeling enabled even if this flag is false.
|
||||
# 即使该标志为 false,服务端仍会启用说话人标注。
|
||||
"speaker_diarization": 0,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start local CAM++, FunASR native realtime WSS, and the browser protocol bridge."""
|
||||
"""启动本地 CAM++、FunASR 原生实时 WSS 服务和浏览器协议桥接层。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -16,8 +16,8 @@ from urllib.request import urlopen
|
|||
from dotenv import load_dotenv
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
# Direct execution from backend/ puts that directory first; prepend the repository
|
||||
# root so backend modules and the local model manifest resolve consistently.
|
||||
# 直接从 backend/ 执行脚本时,该目录会排在搜索路径首位;
|
||||
# 将仓库根目录提前加入,确保后端模块和本地模型清单都能正确导入。
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from backend.model_manifest import (
|
||||
|
|
@ -32,7 +32,7 @@ load_dotenv(PROJECT_ROOT / ".env")
|
|||
|
||||
|
||||
def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
||||
"""Resolve a local ASR or VAD model through the project manifest."""
|
||||
"""通过项目模型清单查找本地 ASR 或 VAD 模型。"""
|
||||
name = requested.strip()
|
||||
configured_path = Path(name)
|
||||
direct_candidates = (
|
||||
|
|
@ -58,7 +58,7 @@ def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
|||
|
||||
|
||||
def local_cam_model(models_dir: Path) -> Path:
|
||||
"""Require a complete CAM++ speaker verification asset from the manifest."""
|
||||
"""从模型清单中确认 CAM++ 说话人验证资源完整。"""
|
||||
manifest = load_manifest()
|
||||
override = os.getenv("CAM_MODEL_PATH", "").strip()
|
||||
if override:
|
||||
|
|
@ -88,7 +88,7 @@ def local_cam_model(models_dir: Path) -> Path:
|
|||
|
||||
|
||||
def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int = 300) -> None:
|
||||
"""Wait for an HTTP child to become ready, reporting early process exit."""
|
||||
"""等待 HTTP 子进程就绪;若进程提前退出则报告错误。"""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
|
|
@ -106,7 +106,7 @@ def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int
|
|||
|
||||
|
||||
def wait_for_tcp(host: str, port: int, process: subprocess.Popen, seconds: int = 300) -> None:
|
||||
"""Wait until FunASR has loaded its models and opened the internal WS socket."""
|
||||
"""等待 FunASR 加载模型并打开内部 WebSocket。"""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
|
|
@ -115,7 +115,7 @@ def wait_for_tcp(host: str, port: int, process: subprocess.Popen, seconds: int =
|
|||
try:
|
||||
with socket.create_connection((host, port), timeout=0.5):
|
||||
pass
|
||||
# Catch an address-in-use failure instead of accepting another process's port.
|
||||
# 遇到端口占用时明确报错,避免误将其他进程的端口视为本服务。
|
||||
time.sleep(0.5)
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
|
|
@ -127,7 +127,7 @@ def wait_for_tcp(host: str, port: int, process: subprocess.Popen, seconds: int =
|
|||
|
||||
|
||||
def stop_child(process: subprocess.Popen | None) -> None:
|
||||
"""Stop a supervised model or WebSocket process on launcher shutdown."""
|
||||
"""启动器关闭时停止受其管理的模型或 WebSocket 进程。"""
|
||||
if process is None or process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
|
|
@ -139,7 +139,7 @@ def stop_child(process: subprocess.Popen | None) -> None:
|
|||
|
||||
|
||||
def main() -> None:
|
||||
"""Require ASR, VAD, and CAM++ before exposing the public WS bridge."""
|
||||
"""确认 ASR、VAD 和 CAM++ 均可用后,再开放公网 WebSocket 桥接服务。"""
|
||||
models_dir = Path(os.getenv("MODEL_DIR", "models"))
|
||||
if not models_dir.is_absolute():
|
||||
models_dir = PROJECT_ROOT / models_dir
|
||||
|
|
@ -170,7 +170,7 @@ def main() -> None:
|
|||
item.strip() for item in env.get("AUXILIARY_PRELOAD_KINDS", "speaker_verification").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
# Native FunASR owns realtime VAD; don't load a duplicate VAD in the CAM++ process.
|
||||
# 实时 VAD 由 FunASR 原生服务负责;不要在 CAM++ 进程中重复加载。
|
||||
preload_kinds.discard("vad")
|
||||
preload_kinds.add("speaker_verification")
|
||||
env.update(
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start the FunASR-backed browser demo."""
|
||||
"""启动由 FunASR 驱动的浏览器演示服务。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
|
|||
|
|
@ -1,222 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
"""在宿主机启动独立的 Qwen3-ASR VLLM 服务。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import shutil
|
||||
import signal
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from urllib.error import URLError
|
||||
from urllib.request import urlopen
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
|
||||
# 启动器自动读取 demo/.env;系统环境变量仍然优先,便于部署平台临时覆盖配置。
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
|
||||
# 将服务端口集中在代码变量中维护,启动时不需要额外传入端口参数;健康检查、
|
||||
# VLLM 子进程命令和就绪提示都使用同一个端口,避免配置不一致导致误判。
|
||||
SERVER_PORT = int(os.getenv("VLLM_PORT", "9950"))
|
||||
|
||||
from backend.model_manifest import load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
def has_model_weights(model_path: Path) -> bool:
|
||||
"""检查 VLLM 加载模型前必须存在的最小本地文件集合。"""
|
||||
if not model_path.is_dir() or not (model_path / "config.json").is_file():
|
||||
return False
|
||||
return any(model_path.rglob("*.safetensors")) or any(model_path.rglob("*.bin"))
|
||||
|
||||
|
||||
def positive_int(value: str) -> int:
|
||||
"""解析启动轮询使用的正整数参数,并拒绝零和负数。"""
|
||||
parsed = int(value)
|
||||
if parsed < 1:
|
||||
raise argparse.ArgumentTypeError("value must be at least 1")
|
||||
return parsed
|
||||
|
||||
|
||||
def non_negative_float(value: str) -> float:
|
||||
"""解析启动轮询间隔,并拒绝会导致逻辑异常的负数。"""
|
||||
parsed = float(value)
|
||||
if parsed < 0:
|
||||
raise argparse.ArgumentTypeError("value must be non-negative")
|
||||
return parsed
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
"""创建宿主机启动参数解析器,默认值允许通过环境变量统一覆盖。"""
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
add_arguments(parser)
|
||||
return parser
|
||||
|
||||
|
||||
def add_arguments(parser: argparse.ArgumentParser) -> None:
|
||||
"""注册模型、网络端点和启动检查循环相关的命令行参数。"""
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default=os.getenv("QWEN3_ASR_MODEL", "default"),
|
||||
help="Model alias, exact model ID, or default",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models-dir",
|
||||
type=Path,
|
||||
default=Path(os.getenv("MODEL_DIR", str(Path(__file__).resolve().parents[1] / "models"))),
|
||||
help="Root directory containing downloaded model files",
|
||||
)
|
||||
parser.add_argument("--host", default=os.getenv("VLLM_HOST", "0.0.0.0"))
|
||||
parser.add_argument(
|
||||
"--display-host",
|
||||
default=os.getenv("VLLM_DISPLAY_HOST", "127.0.0.1"),
|
||||
help="Host name shown in the ready message; does not change the bind address",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--probe-host",
|
||||
default=os.getenv("VLLM_PROBE_HOST", "127.0.0.1"),
|
||||
help="Host used by the startup health probe",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--startup-check-loops",
|
||||
type=positive_int,
|
||||
default=positive_int(os.getenv("VLLM_STARTUP_CHECK_LOOPS", "60")),
|
||||
help="Maximum number of health checks before startup fails",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--startup-check-interval",
|
||||
type=non_negative_float,
|
||||
default=non_negative_float(os.getenv("VLLM_STARTUP_CHECK_INTERVAL_SECONDS", "2")),
|
||||
help="Seconds between startup health checks",
|
||||
)
|
||||
parser.add_argument("--served-model-name", default=os.getenv("VLLM_SERVED_MODEL_NAME"))
|
||||
parser.add_argument(
|
||||
"--gpu-memory-utilization",
|
||||
default=os.getenv("VLLM_GPU_MEMORY_UTILIZATION", "0.3"),
|
||||
)
|
||||
parser.add_argument("--max-model-len", default=os.getenv("VLLM_MAX_MODEL_LEN", "16384"))
|
||||
parser.add_argument("--max-num-seqs", default=os.getenv("VLLM_MAX_NUM_SEQS", "16"))
|
||||
parser.add_argument("--tensor-parallel-size", default=os.getenv("VLLM_TENSOR_PARALLEL_SIZE", "1"))
|
||||
parser.add_argument(
|
||||
"--enforce-eager",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=os.getenv("VLLM_ENFORCE_EAGER", "true").lower() == "true",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--realtime",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=os.getenv("VLLM_REALTIME", "true").lower() == "true",
|
||||
help="Enable the vLLM realtime WebSocket architecture",
|
||||
)
|
||||
|
||||
|
||||
def build_server_command(args: argparse.Namespace, model_id: str, model_path: Path) -> list[str]:
|
||||
"""构造新版 VLLM 原生启动命令,不依赖 qwen-asr-serve。"""
|
||||
# Qwen3-ASR 已由新版 VLLM 原生支持,因此这里调用 vllm serve,避免
|
||||
# qwen-asr-serve 对旧版 VLLM 的固定依赖影响 GB10 部署环境。
|
||||
executable_name = os.getenv("VLLM_EXECUTABLE", "vllm")
|
||||
executable = shutil.which(executable_name)
|
||||
if executable is None:
|
||||
raise RuntimeError(f"{executable_name} was not found; install requirements.txt first")
|
||||
|
||||
served_model_name = args.served_model_name or model_id
|
||||
command = [
|
||||
executable,
|
||||
"serve",
|
||||
str(model_path),
|
||||
"--host",
|
||||
args.host,
|
||||
"--port",
|
||||
str(SERVER_PORT),
|
||||
"--served-model-name",
|
||||
served_model_name,
|
||||
"--gpu-memory-utilization",
|
||||
str(args.gpu_memory_utilization),
|
||||
"--max-model-len",
|
||||
str(args.max_model_len),
|
||||
"--max-num-seqs",
|
||||
str(args.max_num_seqs),
|
||||
"--tensor-parallel-size",
|
||||
str(args.tensor_parallel_size),
|
||||
]
|
||||
if args.realtime:
|
||||
command.extend(["--hf-overrides", '{"architectures":["Qwen3ASRRealtimeGeneration"]}'])
|
||||
if args.enforce_eager:
|
||||
command.append("--enforce-eager")
|
||||
return command
|
||||
|
||||
|
||||
def wait_until_ready(process: subprocess.Popen[bytes], probe_url: str, loops: int, interval: float) -> None:
|
||||
"""按调用方指定的次数和间隔轮询 VLLM 健康接口,直到服务就绪或失败。"""
|
||||
for attempt in range(1, loops + 1):
|
||||
if process.poll() is not None:
|
||||
raise RuntimeError(f"VLLM exited during startup with code {process.returncode}")
|
||||
try:
|
||||
with urlopen(probe_url, timeout=2) as response:
|
||||
if 200 <= response.status < 300:
|
||||
return
|
||||
except (OSError, URLError):
|
||||
pass
|
||||
print(f"Waiting for VLLM startup ({attempt}/{loops})...", flush=True)
|
||||
if attempt < loops:
|
||||
time.sleep(interval)
|
||||
raise TimeoutError(f"VLLM did not become ready after {loops} health checks: {probe_url}")
|
||||
|
||||
|
||||
def stop_process(process: subprocess.Popen[bytes]) -> None:
|
||||
"""向 VLLM 子进程转发优雅停止信号,并在超时后执行兜底清理。"""
|
||||
if process.poll() is not None:
|
||||
return
|
||||
if os.name == "nt":
|
||||
process.send_signal(signal.CTRL_BREAK_EVENT)
|
||||
else:
|
||||
process.send_signal(signal.SIGINT)
|
||||
try:
|
||||
process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.terminate()
|
||||
process.wait(timeout=10)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
"""解析一个本地模型、启动 VLLM,并持续托管宿主机子进程。"""
|
||||
args = build_parser().parse_args()
|
||||
|
||||
manifest = load_manifest()
|
||||
model_id = resolve_model_id(args.model, manifest)
|
||||
model_path = model_directory(model_id, manifest, args.models_dir.resolve())
|
||||
if not has_model_weights(model_path):
|
||||
print(f"Model is missing or incomplete: {model_id} ({model_path})", file=sys.stderr)
|
||||
print("Run scripts/download_models.py for the same model first.", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
command = build_server_command(args, model_id, model_path)
|
||||
probe_url = f"http://{args.probe_host}:{SERVER_PORT}/health"
|
||||
display_url = f"http://{args.display_host}:{SERVER_PORT}"
|
||||
print(f"Starting Qwen3-ASR VLLM service on host: {display_url}", flush=True)
|
||||
print(f"Model: {model_id}", flush=True)
|
||||
print(f"Startup checks: {args.startup_check_loops} x {args.startup_check_interval}s", flush=True)
|
||||
|
||||
process = subprocess.Popen(command)
|
||||
try:
|
||||
wait_until_ready(process, probe_url, args.startup_check_loops, args.startup_check_interval)
|
||||
print(f"VLLM ready: {display_url}", flush=True)
|
||||
print(f"OpenAI endpoint: {display_url}/v1", flush=True)
|
||||
while process.poll() is None:
|
||||
time.sleep(0.5)
|
||||
return int(process.returncode or 0)
|
||||
except (KeyboardInterrupt, TimeoutError, RuntimeError) as exc:
|
||||
print(str(exc), file=sys.stderr)
|
||||
return 1
|
||||
finally:
|
||||
stop_process(process)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
@ -1,76 +0,0 @@
|
|||
"""独立模型清单测试,确保不会导入原项目应用。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from scripts.download_models import fix_camplusplus_config
|
||||
from backend.model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
class ModelManifestTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.manifest = load_manifest()
|
||||
|
||||
def test_default_is_zero_point_six_b_model(self) -> None:
|
||||
self.assertEqual(resolve_model_id("default", self.manifest), "Qwen/Qwen3-ASR-0.6B")
|
||||
|
||||
def test_aliases_resolve_to_individual_models(self) -> None:
|
||||
self.assertEqual(resolve_model_id("1.7b", self.manifest), "Qwen/Qwen3-ASR-1.7B")
|
||||
self.assertEqual(resolve_model_id("0.6b", self.manifest), "Qwen/Qwen3-ASR-0.6B")
|
||||
|
||||
def test_manifest_contains_legacy_and_funasr_asr_models(self) -> None:
|
||||
self.assertTrue({"Qwen/Qwen3-ASR-0.6B", "Qwen/Qwen3-ASR-1.7B"}.issubset(
|
||||
set(self.manifest["models"])
|
||||
))
|
||||
self.assertEqual(
|
||||
resolve_model_id("paraformer-zh-streaming", self.manifest),
|
||||
"iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
)
|
||||
|
||||
def test_manifest_has_auxiliary_runtime_assets(self) -> None:
|
||||
assets = auxiliary_models(self.manifest)
|
||||
self.assertIn("damo/speech_fsmn_vad_zh-cn-16k-common-pytorch", assets)
|
||||
self.assertIn("iic/speech_campplus_speaker-diarization_common", assets)
|
||||
self.assertIn("iic/speech_campplus_sv_zh-cn_16k-common", assets)
|
||||
self.assertIn("Qwen/Qwen3-ForcedAligner-0.6B", assets)
|
||||
|
||||
def test_model_directory_is_under_demo_models(self) -> None:
|
||||
models_dir = Path(__file__).resolve().parents[2] / "models"
|
||||
for model_id in [*self.manifest["models"], *auxiliary_models(self.manifest)]:
|
||||
self.assertTrue(model_directory(model_id, self.manifest, models_dir).is_relative_to(models_dir))
|
||||
|
||||
def test_camplusplus_config_is_rewritten_to_local_assets(self) -> None:
|
||||
"""离线模型包不能继续从 ModelScope 解析 CAM++ 依赖。"""
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
models_dir = Path(temp_dir)
|
||||
config_dir = models_dir / "iic/speech_campplus_speaker-diarization_common"
|
||||
config_dir.mkdir(parents=True)
|
||||
for relative_path in (
|
||||
"damo/speech_campplus_sv_zh-cn_16k-common",
|
||||
"iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
):
|
||||
(models_dir / relative_path).mkdir(parents=True)
|
||||
config = {
|
||||
"model": {
|
||||
"speaker_model": "iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
"change_locator": "damo/speech_campplus_sv_zh-cn_16k-common",
|
||||
"vad_model": "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
}
|
||||
}
|
||||
(config_dir / "configuration.json").write_text(json.dumps(config), encoding="utf-8")
|
||||
|
||||
self.assertTrue(fix_camplusplus_config(models_dir))
|
||||
updated = json.loads((config_dir / "configuration.json").read_text(encoding="utf-8"))
|
||||
self.assertEqual(
|
||||
updated["model"]["speaker_model"],
|
||||
str(models_dir / "iic/speech_campplus_sv_zh-cn_16k-common"),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -1,45 +0,0 @@
|
|||
"""宿主机 VLLM 启动配置测试。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from backend.serve_qwen_legacy import SERVER_PORT, build_parser, build_server_command
|
||||
|
||||
|
||||
class ServeConfigTests(unittest.TestCase):
|
||||
def test_host_and_startup_loop_are_read_from_environment(self) -> None:
|
||||
values = {
|
||||
"VLLM_HOST": "192.168.1.10",
|
||||
"VLLM_DISPLAY_HOST": "asr.local",
|
||||
"VLLM_STARTUP_CHECK_LOOPS": "12",
|
||||
"VLLM_STARTUP_CHECK_INTERVAL_SECONDS": "0.5",
|
||||
}
|
||||
with patch.dict(os.environ, values, clear=False):
|
||||
parser = build_parser()
|
||||
args = parser.parse_args([])
|
||||
|
||||
self.assertEqual(args.host, "192.168.1.10")
|
||||
self.assertEqual(SERVER_PORT, 9950)
|
||||
self.assertNotIn("--port", parser.format_help())
|
||||
self.assertEqual(args.display_host, "asr.local")
|
||||
self.assertEqual(args.startup_check_loops, 12)
|
||||
self.assertEqual(args.startup_check_interval, 0.5)
|
||||
|
||||
@patch("backend.serve_qwen_legacy.shutil.which", return_value="/opt/asr-gb10/bin/vllm")
|
||||
def test_builds_native_vllm_serve_command(self, _which: object) -> None:
|
||||
"""启动器应生成已验证的新版 vllm serve 命令。"""
|
||||
args = build_parser().parse_args([])
|
||||
command = build_server_command(args, "Qwen/Qwen3-ASR-0.6B", Path("/models/Qwen3-ASR-0.6B"))
|
||||
|
||||
self.assertEqual(command[0:3], ["/opt/asr-gb10/bin/vllm", "serve", str(Path("/models/Qwen3-ASR-0.6B"))])
|
||||
self.assertIn("--enforce-eager", command)
|
||||
self.assertIn("--hf-overrides", command)
|
||||
self.assertTrue(any("Qwen3ASRRealtimeGeneration" in item for item in command))
|
||||
self.assertIn("9950", command)
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Serve the unchanged Tencent demo UI and proxy its API to the backend."""
|
||||
"""提供未修改的腾讯演示界面,并将其 API 请求代理到后端。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
@ -25,7 +25,7 @@ HTTP_SESSION = web.AppKey("http_session", ClientSession)
|
|||
|
||||
|
||||
def backend_url(request: web.Request) -> str:
|
||||
"""Keep the original path and query while routing through the backend port."""
|
||||
"""转发到后端端口时保留原始路径和查询参数。"""
|
||||
return f"{BACKEND_BASE_URL}{request.rel_url}"
|
||||
|
||||
|
||||
|
|
@ -36,7 +36,7 @@ async def index_handler(_: web.Request) -> web.FileResponse:
|
|||
|
||||
|
||||
async def api_stop_proxy(request: web.Request) -> web.Response:
|
||||
"""Forward the Tencent page's existing stop request to the WS backend."""
|
||||
"""将腾讯页面现有的停止请求转发给 WebSocket 后端。"""
|
||||
async with request.app[HTTP_SESSION].get(
|
||||
backend_url(request), timeout=ClientTimeout(total=5)
|
||||
) as response:
|
||||
|
|
@ -49,7 +49,7 @@ async def api_stop_proxy(request: web.Request) -> web.Response:
|
|||
|
||||
|
||||
async def websocket_proxy(request: web.Request) -> web.WebSocketResponse:
|
||||
"""Relay text and binary frames without changing the Tencent browser protocol."""
|
||||
"""转发文本帧和二进制帧,不改变腾讯浏览器协议。"""
|
||||
browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30)
|
||||
await browser_ws.prepare(request)
|
||||
try:
|
||||
|
|
@ -104,7 +104,7 @@ async def create_app() -> web.Application:
|
|||
app = web.Application()
|
||||
|
||||
async def lifecycle(application: web.Application):
|
||||
# An unbounded total timeout allows long recordings and slow model loads.
|
||||
# 总超时不设上限,以支持长录音和耗时较长的模型加载。
|
||||
application[HTTP_SESSION] = ClientSession(
|
||||
timeout=ClientTimeout(total=None, connect=10, sock_connect=10, sock_read=None)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
// ===== DOM Elements =====
|
||||
// ===== 页面元素 =====
|
||||
const elEngineModel = document.getElementById('engineModel');
|
||||
const elSpeakerDiarization = document.getElementById('speakerDiarization');
|
||||
const elDiarizationLabel = document.getElementById('diarizationLabel');
|
||||
|
|
@ -19,13 +19,13 @@ const elMicStatus = document.getElementById('micStatus');
|
|||
const elMicTimer = document.getElementById('micTimer');
|
||||
const elMicElapsed = document.getElementById('micElapsed');
|
||||
|
||||
// Input mode tabs
|
||||
// 输入模式选项卡
|
||||
const elTabMic = document.getElementById('tabMic');
|
||||
const elTabFile = document.getElementById('tabFile');
|
||||
const elPanelMic = document.getElementById('panelMic');
|
||||
const elPanelFile = document.getElementById('panelFile');
|
||||
|
||||
// File upload
|
||||
// 文件上传
|
||||
const elAudioFile = document.getElementById('audioFile');
|
||||
const elFileInfo = document.getElementById('fileInfo');
|
||||
const elAudioMeta = document.getElementById('audioMeta');
|
||||
|
|
@ -36,12 +36,12 @@ const elSpeedControl = document.getElementById('speedControl');
|
|||
const elSpeedSlider = document.getElementById('speedSlider');
|
||||
const elSpeedValue = document.getElementById('speedValue');
|
||||
|
||||
// ===== Speaker Diarization Toggle =====
|
||||
// ===== 说话人分离开关 =====
|
||||
elSpeakerDiarization.addEventListener('change', () => {
|
||||
elDiarizationLabel.textContent = elSpeakerDiarization.checked ? '开启' : '关闭';
|
||||
});
|
||||
|
||||
// ===== Log Area =====
|
||||
// ===== 日志区域 =====
|
||||
elBtnClearLog.addEventListener('click', () => { elLogArea.innerHTML = ''; });
|
||||
|
||||
function appendLog(msg) {
|
||||
|
|
@ -57,7 +57,7 @@ function appendLog(msg) {
|
|||
elLogArea.scrollTop = elLogArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== State =====
|
||||
// ===== 状态 =====
|
||||
let ws = null;
|
||||
let sending = false;
|
||||
let stoppingByUser = false;
|
||||
|
|
@ -71,7 +71,7 @@ let micWorklet = null;
|
|||
let micTimerInterval = null;
|
||||
let micStartTime = 0;
|
||||
|
||||
// Input mode (mic / file)
|
||||
// 输入模式(麦克风 / 文件)
|
||||
let inputMode = 'mic';
|
||||
let selectedFile = null;
|
||||
|
||||
|
|
@ -95,7 +95,7 @@ let sentenceMap = {};
|
|||
let speakerOrderMap = {};
|
||||
let speakerOrderCounter = 0;
|
||||
|
||||
// ===== Input Mode Tabs =====
|
||||
// ===== 输入模式选项卡 =====
|
||||
function switchMode(mode) {
|
||||
inputMode = mode;
|
||||
elTabMic.classList.toggle('active', mode === 'mic');
|
||||
|
|
@ -111,7 +111,7 @@ function switchMode(mode) {
|
|||
elTabMic.addEventListener('click', () => switchMode('mic'));
|
||||
elTabFile.addEventListener('click', () => switchMode('file'));
|
||||
|
||||
// ===== File Selection =====
|
||||
// ===== 文件选择 =====
|
||||
elAudioFile.addEventListener('change', (e) => {
|
||||
const file = e.target.files[0];
|
||||
if (!file) return;
|
||||
|
|
@ -135,12 +135,12 @@ elAudioFile.addEventListener('change', (e) => {
|
|||
parseAudioMeta(file);
|
||||
});
|
||||
|
||||
// Speed slider
|
||||
// 语速滑块
|
||||
elSpeedSlider.addEventListener('input', () => {
|
||||
elSpeedValue.textContent = parseFloat(elSpeedSlider.value).toFixed(1) + 'x';
|
||||
});
|
||||
|
||||
// ===== Audio Meta Parsing =====
|
||||
// ===== 音频元信息解析 =====
|
||||
function getFileExt(filename) {
|
||||
const parts = filename.split('.');
|
||||
return parts.length > 1 ? parts[parts.length - 1].toLowerCase() : '';
|
||||
|
|
@ -205,7 +205,7 @@ async function parseAudioMeta(file) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== Copy & Toast =====
|
||||
// ===== 复制与提示 =====
|
||||
function showToast(message, isError) {
|
||||
const toast = document.createElement('div');
|
||||
toast.className = 'copy-toast' + (isError ? ' copy-toast-error' : '');
|
||||
|
|
@ -230,7 +230,7 @@ function handleCopyClick(btn, textEl) {
|
|||
|
||||
elBtnCopyVoiceId.addEventListener('click', () => handleCopyClick(elBtnCopyVoiceId, elVoiceIdDisplay));
|
||||
|
||||
// ===== WAV Export =====
|
||||
// ===== WAV 导出 =====
|
||||
function buildWavBlob(pcmChunks) {
|
||||
let totalLen = 0;
|
||||
for (const c of pcmChunks) totalLen += c.byteLength;
|
||||
|
|
@ -276,7 +276,7 @@ elBtnExportWav.addEventListener('click', () => {
|
|||
showToast('WAV 已导出');
|
||||
});
|
||||
|
||||
// ===== Helpers =====
|
||||
// ===== 辅助函数 =====
|
||||
function formatTime(ms) {
|
||||
const totalSec = Math.floor(ms / 1000);
|
||||
const min = String(Math.floor(totalSec / 60)).padStart(2, '0');
|
||||
|
|
@ -293,7 +293,7 @@ function setStatus(state, text) {
|
|||
elStatusText.textContent = text;
|
||||
}
|
||||
|
||||
// ===== Render: Subtitle (no diarization) =====
|
||||
// ===== 渲染:字幕模式(不启用说话人分离) =====
|
||||
// 每个 sentence_id 对应一个独立气泡:
|
||||
// - 中间态(sentence_type===0):实时更新该气泡文本,末尾加 "..." 表示未定
|
||||
// - 稳态(sentence_type===1):去掉 "..." 并定格,下一个 sentence_id 自动创建新气泡
|
||||
|
|
@ -331,7 +331,7 @@ function renderSubtitle(sentence) {
|
|||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== Render: Speaker Bubble =====
|
||||
// ===== 渲染:说话人气泡 =====
|
||||
function renderBubble(sentence) {
|
||||
const id = 'sent-' + sentence.sentence_id;
|
||||
const speakerId = sentence.speaker_id;
|
||||
|
|
@ -340,7 +340,7 @@ function renderBubble(sentence) {
|
|||
|
||||
if (isUnknown) {
|
||||
const pendingText = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
// Keep unclassified text separate so it is not attributed to the prior speaker.
|
||||
// 将未分类文本单独显示,避免错误归到上一位说话人名下。
|
||||
renderFallbackPendingBubble(id, pendingText, isInterim, sentence);
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
return;
|
||||
|
|
@ -441,7 +441,7 @@ function renderFallbackPendingBubble(id, text, isInterim, sentence) {
|
|||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
}
|
||||
|
||||
// ===== Start Recognition =====
|
||||
// ===== 开始识别 =====
|
||||
elBtnStart.addEventListener('click', () => {
|
||||
if (inputMode === 'file' && !selectedFile) return;
|
||||
startRecognition();
|
||||
|
|
@ -578,7 +578,7 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== Stop =====
|
||||
// ===== 停止 =====
|
||||
elBtnStop.addEventListener('click', () => stopRecognition());
|
||||
|
||||
function stopRecognition() {
|
||||
|
|
@ -622,7 +622,7 @@ function resetControls() {
|
|||
elBtnStop.disabled = true;
|
||||
}
|
||||
|
||||
// ===== Send Audio File =====
|
||||
// ===== 发送音频文件 =====
|
||||
// 按 16KB 切片发送,后端会缓冲成 6400 字节块并按 speed_factor 限流
|
||||
const UPLOAD_CHUNK_SIZE = 16000;
|
||||
async function sendAudioFile(file) {
|
||||
|
|
@ -645,7 +645,7 @@ async function sendAudioFile(file) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== Microphone Capture =====
|
||||
// ===== 麦克风采集 =====
|
||||
async function startMicCapture() {
|
||||
try {
|
||||
micStream = await navigator.mediaDevices.getUserMedia({
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
/* ===== Reset & Base ===== */
|
||||
/* ===== 重置与基础样式 ===== */
|
||||
*,
|
||||
*::before,
|
||||
*::after {
|
||||
|
|
@ -15,7 +15,7 @@ body {
|
|||
min-height: 100vh;
|
||||
}
|
||||
|
||||
/* ===== Two-Column Layout ===== */
|
||||
/* ===== 双栏布局 ===== */
|
||||
.layout {
|
||||
display: flex;
|
||||
height: 100vh;
|
||||
|
|
@ -40,7 +40,7 @@ body {
|
|||
flex-direction: column;
|
||||
}
|
||||
|
||||
/* ===== Header ===== */
|
||||
/* ===== 页头 ===== */
|
||||
header {
|
||||
text-align: center;
|
||||
margin-bottom: 20px;
|
||||
|
|
@ -58,7 +58,7 @@ header h1 {
|
|||
margin-top: 2px;
|
||||
}
|
||||
|
||||
/* ===== Card ===== */
|
||||
/* ===== 卡片 ===== */
|
||||
.card {
|
||||
background: #fff;
|
||||
border-radius: 10px;
|
||||
|
|
@ -80,7 +80,7 @@ header h1 {
|
|||
border-bottom: 1px solid #eee;
|
||||
}
|
||||
|
||||
/* ===== Collapsible Card Header ===== */
|
||||
/* ===== 可折叠卡片标题 ===== */
|
||||
.card-header-collapsible {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
|
|
@ -112,13 +112,13 @@ header h1 {
|
|||
border-top: 1px solid #eee;
|
||||
}
|
||||
|
||||
/* ===== Section Divider ===== */
|
||||
/* ===== 区域分隔线 ===== */
|
||||
.section-divider {
|
||||
border-top: 1px solid #eee;
|
||||
margin: 12px 0;
|
||||
}
|
||||
|
||||
/* ===== Form (left panel) ===== */
|
||||
/* ===== 表单(左侧面板) ===== */
|
||||
.form-stack {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
|
|
@ -185,7 +185,7 @@ input[type="text"][readonly]:focus {
|
|||
box-shadow: none;
|
||||
}
|
||||
|
||||
/* ===== Toggle Switch ===== */
|
||||
/* ===== 开关控件 ===== */
|
||||
.toggle {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
|
|
@ -232,14 +232,14 @@ input[type="text"][readonly]:focus {
|
|||
color: #666;
|
||||
}
|
||||
|
||||
/* ===== Audio Input (left panel) ===== */
|
||||
/* ===== 音频输入(左侧面板) ===== */
|
||||
.audio-input-stack {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
/* Input mode tabs */
|
||||
/* 输入模式选项卡 */
|
||||
.input-mode-tabs {
|
||||
display: flex;
|
||||
gap: 0;
|
||||
|
|
@ -274,7 +274,7 @@ input[type="text"][readonly]:focus {
|
|||
background: #f0f5ff;
|
||||
}
|
||||
|
||||
/* File select */
|
||||
/* 文件选择 */
|
||||
.file-select {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
|
|
@ -292,7 +292,7 @@ input[type="text"][readonly]:focus {
|
|||
font-weight: 500;
|
||||
}
|
||||
|
||||
/* Audio meta info */
|
||||
/* 音频元信息 */
|
||||
.audio-meta {
|
||||
background: #f7f8fa;
|
||||
border-radius: 6px;
|
||||
|
|
@ -316,7 +316,7 @@ input[type="text"][readonly]:focus {
|
|||
font-size: 11px;
|
||||
}
|
||||
|
||||
/* Speed control */
|
||||
/* 语速控制 */
|
||||
.speed-control {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
|
|
@ -353,7 +353,7 @@ input[type="text"][readonly]:focus {
|
|||
cursor: pointer;
|
||||
}
|
||||
|
||||
/* Microphone panel */
|
||||
/* 麦克风面板 */
|
||||
.mic-status {
|
||||
text-align: center;
|
||||
color: #888;
|
||||
|
|
@ -380,7 +380,7 @@ input[type="text"][readonly]:focus {
|
|||
gap: 8px;
|
||||
}
|
||||
|
||||
/* ===== Buttons ===== */
|
||||
/* ===== 按钮 ===== */
|
||||
.btn {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
|
|
@ -446,7 +446,7 @@ input[type="text"][readonly]:focus {
|
|||
background: #e6f7e9;
|
||||
}
|
||||
|
||||
/* Export WAV button */
|
||||
/* 导出 WAV 按钮 */
|
||||
.btn-export {
|
||||
margin-left: 4px;
|
||||
font-size: 11px !important;
|
||||
|
|
@ -458,7 +458,7 @@ input[type="text"][readonly]:focus {
|
|||
cursor: not-allowed;
|
||||
}
|
||||
|
||||
/* ===== Copy Toast ===== */
|
||||
/* ===== 复制提示 ===== */
|
||||
.copy-toast {
|
||||
position: fixed;
|
||||
top: 20px;
|
||||
|
|
@ -487,7 +487,7 @@ input[type="text"][readonly]:focus {
|
|||
background: #ff4d4f;
|
||||
}
|
||||
|
||||
/* ===== Result Section (right panel) ===== */
|
||||
/* ===== 结果区域(右侧面板) ===== */
|
||||
.result-card {
|
||||
flex: 1;
|
||||
min-height: 0;
|
||||
|
|
@ -568,7 +568,7 @@ input[type="text"][readonly]:focus {
|
|||
font-size: 14px;
|
||||
}
|
||||
|
||||
/* ===== Subtitle Mode (no speaker diarization) ===== */
|
||||
/* ===== 字幕模式(不启用说话人分离) ===== */
|
||||
.subtitle-item {
|
||||
display: flex;
|
||||
align-items: flex-start;
|
||||
|
|
@ -614,7 +614,7 @@ input[type="text"][readonly]:focus {
|
|||
opacity: 0.8;
|
||||
}
|
||||
|
||||
/* ===== Bubble Mode (speaker diarization) ===== */
|
||||
/* ===== 气泡模式(启用说话人分离) ===== */
|
||||
.bubble-row {
|
||||
display: flex;
|
||||
margin-bottom: 10px;
|
||||
|
|
@ -690,7 +690,7 @@ input[type="text"][readonly]:focus {
|
|||
font-style: italic;
|
||||
}
|
||||
|
||||
/* Pending text (speaker_id=-1) appended to confirmed bubble */
|
||||
/* 待确认文本(speaker_id=-1)追加到已确认的气泡中 */
|
||||
.pending-text {
|
||||
display: inline;
|
||||
color: #aaa;
|
||||
|
|
@ -710,7 +710,7 @@ input[type="text"][readonly]:focus {
|
|||
font-family: "SF Mono", Menlo, monospace;
|
||||
}
|
||||
|
||||
/* ===== Speaker Colors ===== */
|
||||
/* ===== 说话人颜色 ===== */
|
||||
.speaker-color-0 { background-color: #4a7dff; }
|
||||
.speaker-color-1 { background-color: #52c41a; }
|
||||
.speaker-color-2 { background-color: #faad14; }
|
||||
|
|
@ -723,7 +723,7 @@ input[type="text"][readonly]:focus {
|
|||
.bubble-row.speaker-4 .bubble-body { background: #f9f0ff; }
|
||||
.bubble-row.speaker-5 .bubble-body { background: #e6fffb; }
|
||||
|
||||
/* ===== Responsive ===== */
|
||||
/* ===== 响应式布局 ===== */
|
||||
@media (max-width: 768px) {
|
||||
.layout {
|
||||
flex-direction: column;
|
||||
|
|
@ -741,7 +741,7 @@ input[type="text"][readonly]:focus {
|
|||
}
|
||||
}
|
||||
|
||||
/* ===== Log Panel ===== */
|
||||
/* ===== 日志面板 ===== */
|
||||
.log-card {
|
||||
height: 240px;
|
||||
min-height: 180px;
|
||||
|
|
|
|||
|
|
@ -1,16 +1,6 @@
|
|||
{
|
||||
"default_model": "Qwen/Qwen3-ASR-0.6B",
|
||||
"default_model": "iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
"models": {
|
||||
"Qwen/Qwen3-ASR-1.7B": {
|
||||
"alias": "1.7b",
|
||||
"directory": "Qwen/Qwen3-ASR-1.7B",
|
||||
"description": "Qwen3-ASR 1.7B,GPU 可选大模型"
|
||||
},
|
||||
"Qwen/Qwen3-ASR-0.6B": {
|
||||
"alias": "0.6b",
|
||||
"directory": "Qwen/Qwen3-ASR-0.6B",
|
||||
"description": "Qwen3-ASR 0.6B,GPU 默认轻量模型"
|
||||
},
|
||||
"iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online": {
|
||||
"alias": "paraformer-zh-streaming",
|
||||
"directory": "iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
|
|
@ -125,20 +115,6 @@
|
|||
"transformer_backend.pt"
|
||||
],
|
||||
"min_total_size_bytes": 10000000
|
||||
},
|
||||
"Qwen/Qwen3-ForcedAligner-0.6B": {
|
||||
"alias": "forced-aligner",
|
||||
"directory": "Qwen/Qwen3-ForcedAligner-0.6B",
|
||||
"kind": "forced_aligner",
|
||||
"description": "Qwen3 word-level forced aligner",
|
||||
"required_files": [
|
||||
"config.json"
|
||||
],
|
||||
"any_files": [
|
||||
"*.safetensors",
|
||||
"*.bin"
|
||||
],
|
||||
"min_total_size_bytes": 500000000
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,20 +0,0 @@
|
|||
[build-system]
|
||||
requires = ["setuptools>=68", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "qwen3-asr-vllm-deployment"
|
||||
version = "0.1.0"
|
||||
description = "Standalone Qwen3-ASR model downloader and VLLM deployment"
|
||||
requires-python = ">=3.10,<3.14"
|
||||
dependencies = [
|
||||
"modelscope==1.34.0",
|
||||
"qwen-asr[vllm]==0.0.6",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
qwen3-asr-download = "scripts.download_models:main"
|
||||
qwen3-asr-serve = "backend.serve_qwen_legacy:main"
|
||||
|
||||
[tool.setuptools]
|
||||
packages = ["scripts", "backend", "backend.realtime_websocket"]
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
# Unified dependencies for the frontend, FunASR, CAM++, and Qwen3-ASR vLLM.
|
||||
# The PyTorch and vLLM pins below target the GB10 CUDA 13 deployment.
|
||||
# Update the CUDA index and torch-family pins for a different host.
|
||||
# FunASR 实时后端与前端的统一依赖。
|
||||
# 以下 PyTorch 版本固定值面向 GB10 CUDA 13 部署环境。
|
||||
# 在其他主机上使用时,请相应调整 CUDA 软件源和 PyTorch 相关版本。
|
||||
aiohttp==3.11.11
|
||||
python-dotenv>=1.0
|
||||
numpy>=1.24
|
||||
|
|
@ -14,4 +14,3 @@ websockets>=12,<14
|
|||
torch==2.13.0
|
||||
torchvision==0.28.0
|
||||
torchaudio==2.11.0
|
||||
vllm==0.28.0
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
"""独立 Demo 项目的 VLLM 部署辅助模块。"""
|
||||
"""FunASR 实时演示模型工具。"""
|
||||
|
|
|
|||
|
|
@ -1,12 +1,265 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Download models using the project's shared model manifest."""
|
||||
"""Download FunASR ASR and its configured runtime assets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
# 使用与后端启动器相同的项目本地配置。
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
|
||||
# 同时支持直接执行脚本和 `python -m scripts.download_models` 两种方式,
|
||||
# 下载器只依赖本目录中的清单模块,不耦合原项目的包路径。
|
||||
try:
|
||||
from .download_models_qwen_legacy import main
|
||||
from .model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
except ImportError:
|
||||
from download_models_qwen_legacy import main
|
||||
from model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
def has_model_weights(model_path: Path) -> bool:
|
||||
"""检查流式 ASR 模型所需的最少本地文件。"""
|
||||
if not model_path.is_dir():
|
||||
return False
|
||||
if not (model_path / "config.json").is_file():
|
||||
return False
|
||||
return any(model_path.rglob("*.safetensors")) or any(model_path.rglob("*.bin"))
|
||||
|
||||
|
||||
def is_model_ready(model_path: Path, config: dict[str, object]) -> bool:
|
||||
"""根据模型清单中的专属文件规则检查 ASR 或辅助资产是否完整。"""
|
||||
if not model_path.is_dir():
|
||||
return False
|
||||
required_files = config.get("required_files", [])
|
||||
if isinstance(required_files, list):
|
||||
for relative_path in required_files:
|
||||
if not (model_path / str(relative_path)).is_file():
|
||||
return False
|
||||
|
||||
any_files = config.get("any_files", [])
|
||||
if isinstance(any_files, list) and any_files:
|
||||
if not any(
|
||||
file_path.is_file()
|
||||
for pattern in any_files
|
||||
for file_path in model_path.rglob(str(pattern))
|
||||
):
|
||||
return False
|
||||
|
||||
minimum_size_value = config.get("min_total_size_bytes", 0)
|
||||
# 模型清单使用 object 表示不同类型的资产字段,因此在转换为整数前必须
|
||||
# 先收窄类型,避免不合法的清单值在运行时触发难以定位的类型异常。
|
||||
minimum_size = (
|
||||
int(minimum_size_value)
|
||||
if isinstance(minimum_size_value, (int, str))
|
||||
else 0
|
||||
)
|
||||
if minimum_size:
|
||||
total_size = sum(file_path.stat().st_size for file_path in model_path.rglob("*") if file_path.is_file())
|
||||
if total_size < minimum_size:
|
||||
return False
|
||||
if required_files or any_files or minimum_size:
|
||||
return True
|
||||
return has_model_weights(model_path)
|
||||
|
||||
|
||||
def download_model(
|
||||
model_id: str,
|
||||
model_path: Path,
|
||||
cache_dir: Path | None,
|
||||
revision: str | None,
|
||||
) -> None:
|
||||
"""通过 ModelScope 下载一个指定资产,且不导入原项目应用代码。"""
|
||||
# 延迟导入较重的 ModelScope 模块,使本地资源检查无需安装下载依赖。
|
||||
try:
|
||||
from modelscope.hub.snapshot_download import snapshot_download
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"ModelScope is required for downloading; install requirements.txt first"
|
||||
) from exc
|
||||
|
||||
model_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
cache_path: str | None = None
|
||||
if cache_dir is not None:
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
cache_path = str(cache_dir)
|
||||
print(f"Downloading model asset: {model_id}")
|
||||
print(f"Local directory: {model_path}")
|
||||
# 使用显式关键字参数而不是 **dict,既便于 Pylance 推断 ModelScope 的真实
|
||||
# 参数类型,也避免动态字典被误判为其它无关参数的类型签名。
|
||||
snapshot_download(
|
||||
model_id,
|
||||
revision=revision,
|
||||
cache_dir=cache_path,
|
||||
local_dir=str(model_path),
|
||||
)
|
||||
|
||||
|
||||
def fix_camplusplus_config(models_dir: Path) -> bool:
|
||||
"""将 CAM++ 依赖模型 ID 改写为本地路径,确保服务可以离线启动。
|
||||
|
||||
聚类流水线会在 ``configuration.json`` 中保存多个 ModelScope 模型 ID。
|
||||
如果不改写这些 ID,即使所有文件已经下载完整,辅助服务在无网络环境
|
||||
启动时仍可能再次访问 ModelScope 获取依赖。
|
||||
"""
|
||||
config_file = models_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json"
|
||||
if not config_file.is_file():
|
||||
return False
|
||||
|
||||
replacements = {
|
||||
"damo/speech_campplus_sv_zh-cn_16k-common": models_dir / "damo/speech_campplus_sv_zh-cn_16k-common",
|
||||
"iic/speech_campplus_sv_zh-cn_16k-common": models_dir / "iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn_16k-common",
|
||||
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn-16k-common",
|
||||
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": models_dir / "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
}
|
||||
try:
|
||||
config = json.loads(config_file.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
print(f"Unable to read CAM++ configuration: {exc}")
|
||||
return False
|
||||
|
||||
raw_model_config = config.get("model")
|
||||
if not isinstance(raw_model_config, dict):
|
||||
return False
|
||||
model_config: dict[str, object] = {
|
||||
str(key): value for key, value in raw_model_config.items()
|
||||
}
|
||||
modified = False
|
||||
for key in ("speaker_model", "change_locator", "vad_model"):
|
||||
old_value = model_config.get(key)
|
||||
local_path = replacements.get(old_value) if isinstance(old_value, str) else None
|
||||
if local_path is not None and local_path.exists():
|
||||
model_config[key] = str(local_path)
|
||||
modified = True
|
||||
if not modified:
|
||||
return False
|
||||
config["model"] = model_config
|
||||
config_file.write_text(json.dumps(config, indent=4, ensure_ascii=False) + "\n", encoding="utf-8")
|
||||
return True
|
||||
|
||||
|
||||
def main() -> int:
|
||||
"""检查或下载已配置的 ASR 模型及运行时资源。"""
|
||||
# 将所选资源下载到共享的本地模型目录,以支持离线启动。
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default=os.getenv("FUNASR_ASR_MODEL", "paraformer-zh-streaming"),
|
||||
help="ASR model alias, exact model ID, or default",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models-dir",
|
||||
type=Path,
|
||||
default=Path(os.getenv("MODEL_DIR", str(Path(__file__).resolve().parents[1] / "models"))),
|
||||
help="Root directory for local model files",
|
||||
)
|
||||
model_scope_cache = os.getenv("MODELSCOPE_CACHE")
|
||||
parser.add_argument(
|
||||
"--cache-dir",
|
||||
type=Path,
|
||||
default=Path(model_scope_cache) if model_scope_cache else None,
|
||||
help="Optional ModelScope cache directory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--check-only",
|
||||
action="store_true",
|
||||
help="Only check selected assets; do not download",
|
||||
)
|
||||
auxiliary_group = parser.add_mutually_exclusive_group()
|
||||
auxiliary_group.add_argument(
|
||||
"--skip-auxiliary",
|
||||
action="store_true",
|
||||
help="Only download/check the selected ASR model",
|
||||
)
|
||||
auxiliary_group.add_argument(
|
||||
"--auxiliary-only",
|
||||
action="store_true",
|
||||
help="Only download/check configured auxiliary assets",
|
||||
)
|
||||
auxiliary_group.add_argument(
|
||||
"--funasr-runtime",
|
||||
action="store_true",
|
||||
help="Download/check the configured ASR model, FSMN-VAD, CAM++, and punctuation",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
manifest = load_manifest()
|
||||
models_dir = args.models_dir
|
||||
if not models_dir.is_absolute():
|
||||
models_dir = Path(__file__).resolve().parents[1] / models_dir
|
||||
models_dir = models_dir.resolve()
|
||||
cache_dir = args.cache_dir.resolve() if args.cache_dir else None
|
||||
selected_assets: list[tuple[str, dict[str, object]]] = []
|
||||
if args.funasr_runtime:
|
||||
# 标点模型在运行时可选,但下载后可获得完整的本地输出。
|
||||
asr_id = resolve_model_id(args.model, manifest)
|
||||
selected_assets.append((asr_id, manifest["models"][asr_id]))
|
||||
assets = auxiliary_models(manifest)
|
||||
vad_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "vad"
|
||||
)
|
||||
cam_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "speaker_verification"
|
||||
and model_id.startswith("iic/")
|
||||
)
|
||||
punctuation_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "punctuation"
|
||||
)
|
||||
selected_assets.extend(
|
||||
(model_id, assets[model_id])
|
||||
for model_id in (vad_id, cam_id, punctuation_id)
|
||||
)
|
||||
else:
|
||||
if not args.auxiliary_only:
|
||||
model_id = resolve_model_id(args.model, manifest)
|
||||
selected_assets.append((model_id, manifest["models"][model_id]))
|
||||
if not args.skip_auxiliary:
|
||||
selected_assets.extend(auxiliary_models(manifest).items())
|
||||
|
||||
missing: list[tuple[str, Path, dict[str, object]]] = []
|
||||
for model_id, config in selected_assets:
|
||||
model_path = model_directory(model_id, manifest, models_dir)
|
||||
if is_model_ready(model_path, config):
|
||||
print(f"Model asset is ready: {model_id}")
|
||||
else:
|
||||
missing.append((model_id, model_path, config))
|
||||
|
||||
if not missing:
|
||||
# 即使资产已经存在,也要重新执行一次离线配置修正;这样从其它主机
|
||||
# 复制过来的模型包也能在启动辅助服务前自动完成本地路径修复。
|
||||
if fix_camplusplus_config(models_dir):
|
||||
print("CAM++ configuration updated for offline local model paths")
|
||||
print(f"All selected model assets are ready: {len(selected_assets)}")
|
||||
return 0
|
||||
if args.check_only:
|
||||
for model_id, model_path, _ in missing:
|
||||
print(f"Model asset is missing or incomplete: {model_id} ({model_path})")
|
||||
return 1
|
||||
|
||||
failed: list[str] = []
|
||||
for model_id, model_path, config in missing:
|
||||
try:
|
||||
revision = str(config.get("revision") or "") or None
|
||||
download_model(model_id, model_path, cache_dir, revision)
|
||||
if not is_model_ready(model_path, config):
|
||||
print(f"Download finished but model asset is incomplete: {model_path}")
|
||||
failed.append(model_id)
|
||||
else:
|
||||
print(f"Model asset is ready: {model_id}")
|
||||
except Exception as exc:
|
||||
print(f"Download failed: {model_id}: {exc}")
|
||||
failed.append(model_id)
|
||||
if not failed and fix_camplusplus_config(models_dir):
|
||||
print("CAM++ configuration updated for offline local model paths")
|
||||
return 1 if failed else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -1,261 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
"""为独立服务部署下载 ASR 和配套辅助模型。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# 同时支持直接执行脚本和 `python -m scripts.download_models` 两种方式,
|
||||
# 下载器只依赖本目录中的清单模块,不耦合原项目的包路径。
|
||||
try:
|
||||
from .model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
except ImportError:
|
||||
from model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
def has_model_weights(model_path: Path) -> bool:
|
||||
"""检查 VLLM 加载 ASR 模型前必须存在的最小本地文件集合。"""
|
||||
if not model_path.is_dir():
|
||||
return False
|
||||
if not (model_path / "config.json").is_file():
|
||||
return False
|
||||
return any(model_path.rglob("*.safetensors")) or any(model_path.rglob("*.bin"))
|
||||
|
||||
|
||||
def is_model_ready(model_path: Path, config: dict[str, object]) -> bool:
|
||||
"""根据模型清单中的专属文件规则检查 ASR 或辅助资产是否完整。"""
|
||||
if not model_path.is_dir():
|
||||
return False
|
||||
required_files = config.get("required_files", [])
|
||||
if isinstance(required_files, list):
|
||||
for relative_path in required_files:
|
||||
if not (model_path / str(relative_path)).is_file():
|
||||
return False
|
||||
|
||||
any_files = config.get("any_files", [])
|
||||
if isinstance(any_files, list) and any_files:
|
||||
if not any(
|
||||
file_path.is_file()
|
||||
for pattern in any_files
|
||||
for file_path in model_path.rglob(str(pattern))
|
||||
):
|
||||
return False
|
||||
|
||||
minimum_size_value = config.get("min_total_size_bytes", 0)
|
||||
# 模型清单使用 object 表示不同类型的资产字段,因此在转换为整数前必须
|
||||
# 先收窄类型,避免不合法的清单值在运行时触发难以定位的类型异常。
|
||||
minimum_size = (
|
||||
int(minimum_size_value)
|
||||
if isinstance(minimum_size_value, (int, str))
|
||||
else 0
|
||||
)
|
||||
if minimum_size:
|
||||
total_size = sum(file_path.stat().st_size for file_path in model_path.rglob("*") if file_path.is_file())
|
||||
if total_size < minimum_size:
|
||||
return False
|
||||
if required_files or any_files or minimum_size:
|
||||
return True
|
||||
return has_model_weights(model_path)
|
||||
|
||||
|
||||
def download_model(
|
||||
model_id: str,
|
||||
model_path: Path,
|
||||
cache_dir: Path | None,
|
||||
revision: str | None,
|
||||
) -> None:
|
||||
"""通过 ModelScope 下载一个指定资产,且不导入原项目应用代码。"""
|
||||
# 延迟导入 ModelScope,使模型清单检查和单元测试无需安装重量级依赖。
|
||||
try:
|
||||
from modelscope.hub.snapshot_download import snapshot_download
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"ModelScope is required for downloading; install requirements.txt first"
|
||||
) from exc
|
||||
|
||||
model_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
cache_path: str | None = None
|
||||
if cache_dir is not None:
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
cache_path = str(cache_dir)
|
||||
print(f"Downloading model asset: {model_id}")
|
||||
print(f"Local directory: {model_path}")
|
||||
# 使用显式关键字参数而不是 **dict,既便于 Pylance 推断 ModelScope 的真实
|
||||
# 参数类型,也避免动态字典被误判为其它无关参数的类型签名。
|
||||
snapshot_download(
|
||||
model_id,
|
||||
revision=revision,
|
||||
cache_dir=cache_path,
|
||||
local_dir=str(model_path),
|
||||
)
|
||||
|
||||
|
||||
def fix_camplusplus_config(models_dir: Path) -> bool:
|
||||
"""将 CAM++ 依赖模型 ID 改写为本地路径,确保服务可以离线启动。
|
||||
|
||||
聚类流水线会在 ``configuration.json`` 中保存多个 ModelScope 模型 ID。
|
||||
如果不改写这些 ID,即使所有文件已经下载完整,辅助服务在无网络环境
|
||||
启动时仍可能再次访问 ModelScope 获取依赖。
|
||||
"""
|
||||
config_file = models_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json"
|
||||
if not config_file.is_file():
|
||||
return False
|
||||
|
||||
replacements = {
|
||||
"damo/speech_campplus_sv_zh-cn_16k-common": models_dir / "damo/speech_campplus_sv_zh-cn_16k-common",
|
||||
"iic/speech_campplus_sv_zh-cn_16k-common": models_dir / "iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn_16k-common",
|
||||
"damo/speech_campplus-transformer_scl_zh-cn-16k-common": models_dir / "damo/speech_campplus-transformer_scl_zh-cn-16k-common",
|
||||
"damo/speech_fsmn_vad_zh-cn-16k-common-pytorch": models_dir / "damo/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
}
|
||||
try:
|
||||
config = json.loads(config_file.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
print(f"Unable to read CAM++ configuration: {exc}")
|
||||
return False
|
||||
|
||||
raw_model_config = config.get("model")
|
||||
if not isinstance(raw_model_config, dict):
|
||||
return False
|
||||
model_config: dict[str, object] = {
|
||||
str(key): value for key, value in raw_model_config.items()
|
||||
}
|
||||
modified = False
|
||||
for key in ("speaker_model", "change_locator", "vad_model"):
|
||||
old_value = model_config.get(key)
|
||||
local_path = replacements.get(old_value) if isinstance(old_value, str) else None
|
||||
if local_path is not None and local_path.exists():
|
||||
model_config[key] = str(local_path)
|
||||
modified = True
|
||||
if not modified:
|
||||
return False
|
||||
config["model"] = model_config
|
||||
config_file.write_text(json.dumps(config, indent=4, ensure_ascii=False) + "\n", encoding="utf-8")
|
||||
return True
|
||||
|
||||
|
||||
def main() -> int:
|
||||
"""检查或下载 ASR 模型及辅助运行时所需的全部资产。"""
|
||||
# 保持当前部署项目与原项目模型规划器完全独立,同时将孤立服务需要的
|
||||
# 模型统一准备到本地,方便后续在服务器上离线启动多个常驻服务。
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default=os.getenv("QWEN3_ASR_MODEL", "default"),
|
||||
help="ASR model alias (1.7b/0.6b), exact model ID, or default",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--models-dir",
|
||||
type=Path,
|
||||
default=Path(os.getenv("MODEL_DIR", str(Path(__file__).resolve().parents[1] / "models"))),
|
||||
help="Root directory for local model files",
|
||||
)
|
||||
model_scope_cache = os.getenv("MODELSCOPE_CACHE")
|
||||
parser.add_argument(
|
||||
"--cache-dir",
|
||||
type=Path,
|
||||
default=Path(model_scope_cache) if model_scope_cache else None,
|
||||
help="Optional ModelScope cache directory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--check-only",
|
||||
action="store_true",
|
||||
help="Only check selected assets; do not download",
|
||||
)
|
||||
auxiliary_group = parser.add_mutually_exclusive_group()
|
||||
auxiliary_group.add_argument(
|
||||
"--skip-auxiliary",
|
||||
action="store_true",
|
||||
help="Only download/check the selected ASR model",
|
||||
)
|
||||
auxiliary_group.add_argument(
|
||||
"--auxiliary-only",
|
||||
action="store_true",
|
||||
help="Only download/check VAD, speaker, punctuation, diarization, and aligner assets",
|
||||
)
|
||||
auxiliary_group.add_argument(
|
||||
"--funasr-runtime",
|
||||
action="store_true",
|
||||
help="Download/check streaming Paraformer, FSMN-VAD, CAM++, and punctuation",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
manifest = load_manifest()
|
||||
models_dir = args.models_dir
|
||||
if not models_dir.is_absolute():
|
||||
models_dir = Path(__file__).resolve().parents[1] / models_dir
|
||||
models_dir = models_dir.resolve()
|
||||
cache_dir = args.cache_dir.resolve() if args.cache_dir else None
|
||||
selected_assets: list[tuple[str, dict[str, object]]] = []
|
||||
if args.funasr_runtime:
|
||||
# Punctuation is optional at runtime, but included for complete local output.
|
||||
asr_id = resolve_model_id("paraformer-zh-streaming", manifest)
|
||||
selected_assets.append((asr_id, manifest["models"][asr_id]))
|
||||
assets = auxiliary_models(manifest)
|
||||
vad_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "vad"
|
||||
)
|
||||
cam_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "speaker_verification"
|
||||
and model_id.startswith("iic/")
|
||||
)
|
||||
punctuation_id = next(
|
||||
model_id for model_id, config in assets.items()
|
||||
if config.get("kind") == "punctuation"
|
||||
)
|
||||
selected_assets.extend(
|
||||
(model_id, assets[model_id])
|
||||
for model_id in (vad_id, cam_id, punctuation_id)
|
||||
)
|
||||
else:
|
||||
if not args.auxiliary_only:
|
||||
model_id = resolve_model_id(args.model, manifest)
|
||||
selected_assets.append((model_id, manifest["models"][model_id]))
|
||||
if not args.skip_auxiliary:
|
||||
selected_assets.extend(auxiliary_models(manifest).items())
|
||||
|
||||
missing: list[tuple[str, Path, dict[str, object]]] = []
|
||||
for model_id, config in selected_assets:
|
||||
model_path = model_directory(model_id, manifest, models_dir)
|
||||
if is_model_ready(model_path, config):
|
||||
print(f"Model asset is ready: {model_id}")
|
||||
else:
|
||||
missing.append((model_id, model_path, config))
|
||||
|
||||
if not missing:
|
||||
# 即使资产已经存在,也要重新执行一次离线配置修正;这样从其它主机
|
||||
# 复制过来的模型包也能在启动辅助服务前自动完成本地路径修复。
|
||||
if fix_camplusplus_config(models_dir):
|
||||
print("CAM++ configuration updated for offline local model paths")
|
||||
print(f"All selected model assets are ready: {len(selected_assets)}")
|
||||
return 0
|
||||
if args.check_only:
|
||||
for model_id, model_path, _ in missing:
|
||||
print(f"Model asset is missing or incomplete: {model_id} ({model_path})")
|
||||
return 1
|
||||
|
||||
failed: list[str] = []
|
||||
for model_id, model_path, config in missing:
|
||||
try:
|
||||
revision = str(config.get("revision") or "") or None
|
||||
download_model(model_id, model_path, cache_dir, revision)
|
||||
if not is_model_ready(model_path, config):
|
||||
print(f"Download finished but model asset is incomplete: {model_path}")
|
||||
failed.append(model_id)
|
||||
else:
|
||||
print(f"Model asset is ready: {model_id}")
|
||||
except Exception as exc:
|
||||
print(f"Download failed: {model_id}: {exc}")
|
||||
failed.append(model_id)
|
||||
if not failed and fix_camplusplus_config(models_dir):
|
||||
print("CAM++ configuration updated for offline local model paths")
|
||||
return 1 if failed else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
@ -1,10 +1,10 @@
|
|||
"""Compatibility exports for shared model download scripts."""
|
||||
"""供共享模型下载脚本使用的兼容性导出。"""
|
||||
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
# Direct script execution puts scripts/ first on sys.path; add the repository root
|
||||
# so the shared backend manifest remains importable in both launch modes.
|
||||
# 直接执行脚本时,scripts/ 会排在 sys.path 首位;将仓库根目录
|
||||
# 加入搜索路径,确保两种启动方式都能导入后端共享模型清单。
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
|
|
|||
Loading…
Reference in New Issue