Compare commits
No commits in common. "5c58759099b13f63cbe893d44e0a623ed77fa19e" and "525d8060f59faa2caaccdbd57da67d5438e582b4" have entirely different histories.
5c58759099
...
525d8060f5
85
.env.example
85
.env.example
|
|
@ -1,41 +1,50 @@
|
|||
# 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_PATH=iic/speech_campplus_sv_zh-cn_16k-common
|
||||
# 选择已经下载到 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
|
||||
|
||||
# FunASR realtime engine
|
||||
FUNASR_ASR_MODEL=paraformer-zh-streaming
|
||||
FUNASR_VAD_MODEL=fsmn-vad
|
||||
FUNASR_DEVICE=cuda:0
|
||||
FUNASR_VAD_DEVICE=cpu
|
||||
FUNASR_CHUNK_SIZE=0,10,5
|
||||
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_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
|
||||
# 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
|
||||
|
||||
# CAM++ is required and started by the backend launcher.
|
||||
# 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
|
||||
|
||||
# Standalone frontend and browser-facing backend origin.
|
||||
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
|
||||
|
||||
WEB_HOST=0.0.0.0
|
||||
WEB_PORT=8082
|
||||
WEB_DISPLAY_HOST=127.0.0.1
|
||||
# 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 +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 +0,0 @@
|
|||
# FunASR realtime browser demo
|
||||
|
||||
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.
|
||||
|
||||
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.
|
||||
|
||||
## Model directories
|
||||
|
||||
Put assets under `models/` or set `MODEL_DIR` in `.env`. The default names resolve to these local directories:
|
||||
|
||||
- `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)
|
||||
|
||||
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.
|
||||
|
||||
## 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:
|
||||
|
||||
~~~powershell
|
||||
python -m pip install -r requirements.txt
|
||||
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}`.
|
||||
|
||||
`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.
|
||||
|
||||
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++.
|
||||
163
README.md
163
README.md
|
|
@ -1,32 +1,151 @@
|
|||
# FunASR realtime ASR demo
|
||||
# Qwen3-ASR VLLM 独立部署项目
|
||||
|
||||
The project keeps the frontend and backend in separate directories. The frontend serves the original Tencent demo assets from frontend/static/ unchanged.
|
||||
本目录是后续实时 ASR 功能验证使用的独立模型服务项目。
|
||||
|
||||
## Project layout
|
||||
它不导入、不启动、也不调用仓库根目录下原项目的 `app/` 代码。模型下载、VLLM 启动、配置和服务验证都在本目录内完成。后续验证 demo 只需要调用这里提供的 VLLM OpenAI 兼容接口。
|
||||
|
||||
- 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
|
||||
默认下载一个 ASR 模型和独立辅助模型运行服务所需的全部模型资产:
|
||||
|
||||
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.
|
||||
```text
|
||||
Qwen/Qwen3-ASR-0.6B
|
||||
```
|
||||
|
||||
~~~powershell
|
||||
python -m pip install -r requirements.txt
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.funasr.example .env }
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
ASR 如需使用大模型,可显式选择 1.7B:
|
||||
|
||||
## Start
|
||||
```text
|
||||
Qwen/Qwen3-ASR-1.7B
|
||||
```
|
||||
|
||||
Run the backend and frontend in separate terminals from the project root:
|
||||
辅助资产包括 VAD、CAM++ 分离、配置声纹、实时声纹、CAM++ Transformer 和 Qwen3 ForcedAligner。它们不会由 `qwen-asr-serve` 启动,而是由独立 Python 辅助模型服务加载。
|
||||
|
||||
~~~powershell
|
||||
python backend/run_backend.py
|
||||
python frontend/run_frontend.py
|
||||
~~~
|
||||
不会下载另一个未选择的 ASR 模型;辅助模型资产会随默认部署包下载,供独立 Python 运行服务预加载。
|
||||
|
||||
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.
|
||||
## 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-download.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-deploy.txt
|
||||
```
|
||||
|
||||
先复制并按服务器实际路径修改 `.env`,启动器会自动读取该文件:
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
```
|
||||
|
||||
默认启动 `Qwen/Qwen3-ASR-0.6B`,监听地址为 `0.0.0.0:9950`:
|
||||
|
||||
```bash
|
||||
python scripts/serve.py
|
||||
```
|
||||
|
||||
模型和启动检查循环可以通过命令行或环境变量传入;端口统一在 `scripts/serve.py` 的 `SERVER_PORT` 变量中维护:
|
||||
|
||||
```bash
|
||||
QWEN3_ASR_MODEL=0.6b VLLM_STARTUP_CHECK_LOOPS=120 \
|
||||
VLLM_STARTUP_CHECK_INTERVAL_SECONDS=2 python scripts/serve.py
|
||||
```
|
||||
|
||||
如果使用 `1.7B`,下载和启动必须指定同一个模型:
|
||||
|
||||
```powershell
|
||||
python scripts\download_models.py --model 1.7b
|
||||
python -m scripts.serve --model 1.7b
|
||||
```
|
||||
|
||||
启动器会按 `VLLM_STARTUP_CHECK_LOOPS` 次数轮询 `/health`,每次间隔由 `VLLM_STARTUP_CHECK_INTERVAL_SECONDS` 指定。服务端口、健康检查端口和就绪提示统一使用 `scripts/serve.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-auxiliary.txt
|
||||
python scripts\auxiliary_server.py
|
||||
```
|
||||
|
||||
辅助服务默认监听 `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 和全部辅助模型资产。
|
||||
- `scripts/serve.py`:读取 `.env`,解析模型选择、宿主机参数并启动新版 `vllm serve`。
|
||||
- `requirements-deploy.txt`:安装宿主机部署所需的官方 Qwen3-ASR VLLM 依赖。
|
||||
- `tests/`:只验证本项目自己的模型清单和选择逻辑,不依赖原项目。
|
||||
|
||||
模型服务就绪后,新的实时 ASR demo 放在同级 `demo` 项目中继续开发,但不得通过 Python import 或 HTTP/WebSocket 调用原项目服务。
|
||||
|
|
|
|||
|
|
@ -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 +0,0 @@
|
|||
"""Backend services and the realtime ASR protocol adapter."""
|
||||
|
|
@ -1,79 +0,0 @@
|
|||
"""独立 VLLM 部署项目的模型清单与路径解析辅助函数。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
MANIFEST_PATH = PROJECT_ROOT / "model_manifest.json"
|
||||
|
||||
|
||||
def load_manifest(path: Path = MANIFEST_PATH) -> dict[str, Any]:
|
||||
"""读取本项目自己的模型清单,整个过程不导入原项目代码。"""
|
||||
with path.open("r", encoding="utf-8") as manifest_file:
|
||||
manifest = json.load(manifest_file)
|
||||
if not isinstance(manifest.get("models"), dict) or not manifest["models"]:
|
||||
raise ValueError("model_manifest.json 必须包含非空的 models 对象")
|
||||
return manifest
|
||||
|
||||
|
||||
def resolve_model_id(model: str | None, manifest: dict[str, Any]) -> str:
|
||||
"""将默认值、短别名或完整模型 ID 解析为一个 ASR 模型。"""
|
||||
models = manifest["models"]
|
||||
requested = (model or "default").strip()
|
||||
if requested == "default":
|
||||
requested = str(manifest["default_model"])
|
||||
|
||||
if requested in models:
|
||||
return requested
|
||||
|
||||
for model_id, config in models.items():
|
||||
if requested.lower() == str(config.get("alias", "")).lower():
|
||||
return model_id
|
||||
raise ValueError(f"不支持的 ASR 模型 '{model}',可选模型:{', '.join(models)}")
|
||||
|
||||
|
||||
def model_directory(model_id: str, manifest: dict[str, Any], models_dir: Path) -> Path:
|
||||
"""根据清单返回 ASR 或辅助模型实际使用的本地目录。"""
|
||||
config = manifest.get("models", {}).get(model_id)
|
||||
if config is None:
|
||||
config = manifest.get("auxiliary_models", {}).get(model_id)
|
||||
if not isinstance(config, dict) or not config.get("directory"):
|
||||
raise ValueError(f"模型 '{model_id}' 在清单中没有配置本地目录")
|
||||
return models_dir / str(config["directory"])
|
||||
|
||||
|
||||
def auxiliary_models(manifest: dict[str, Any]) -> dict[str, dict[str, Any]]:
|
||||
"""返回可独立部署的 VAD、说话人和对齐模型资产。"""
|
||||
models = manifest.get("auxiliary_models", {})
|
||||
if not isinstance(models, dict):
|
||||
raise ValueError("model_manifest.json 的 auxiliary_models 必须是对象")
|
||||
return {str(model_id): config for model_id, config in models.items() if isinstance(config, dict)}
|
||||
|
||||
def resolve_auxiliary_model_id(
|
||||
model: str | None,
|
||||
manifest: dict[str, Any],
|
||||
kind: str | None = None,
|
||||
) -> str:
|
||||
"""Resolve a configured auxiliary asset by ID, alias, or FunASR alias."""
|
||||
assets = auxiliary_models(manifest)
|
||||
requested = (model or "").strip()
|
||||
if requested in assets and (kind is None or assets[requested].get("kind") == kind):
|
||||
return requested
|
||||
|
||||
for model_id, config in assets.items():
|
||||
if kind is not None and config.get("kind") != kind:
|
||||
continue
|
||||
aliases = [config.get("alias"), config.get("funasr_alias")]
|
||||
if requested.lower() in {
|
||||
str(alias).lower() for alias in aliases if isinstance(alias, str)
|
||||
}:
|
||||
return model_id
|
||||
expected = ", ".join(
|
||||
model_id for model_id, config in assets.items()
|
||||
if kind is None or config.get("kind") == kind
|
||||
)
|
||||
raise ValueError(f"Unsupported auxiliary model '{model}'; configured assets: {expected}")
|
||||
|
|
@ -1 +0,0 @@
|
|||
"""FunASR realtime browser demo package."""
|
||||
|
|
@ -1,393 +0,0 @@
|
|||
"""FunASR realtime engine migrated into the demo project.
|
||||
|
||||
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
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FunASRServiceConfig:
|
||||
"""Configuration for the in-process FunASR models."""
|
||||
|
||||
model: str = "paraformer-zh-streaming"
|
||||
vad_model: str = "fsmn-vad"
|
||||
device: str = "cuda:0"
|
||||
vad_device: str = "cpu"
|
||||
sample_rate: int = 16000
|
||||
vad_chunk_ms: int = 200
|
||||
chunk_size: tuple[int, int, int] = (0, 10, 5)
|
||||
encoder_chunk_look_back: int = 4
|
||||
decoder_chunk_look_back: int = 1
|
||||
max_segment_sec: float = 30.0
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "FunASRServiceConfig":
|
||||
"""Read model selection from environment without changing WS fields."""
|
||||
chunk_text = os.getenv("FUNASR_CHUNK_SIZE", "0,10,5")
|
||||
try:
|
||||
values = tuple(int(value.strip()) for value in chunk_text.split(","))
|
||||
chunk_size = values if len(values) == 3 else cls.chunk_size
|
||||
except ValueError:
|
||||
chunk_size = cls.chunk_size
|
||||
return cls(
|
||||
model=os.getenv("FUNASR_ASR_MODEL", cls.model),
|
||||
vad_model=os.getenv("FUNASR_VAD_MODEL", cls.vad_model),
|
||||
device=os.getenv("FUNASR_DEVICE", cls.device),
|
||||
vad_device=os.getenv("FUNASR_VAD_DEVICE", cls.vad_device),
|
||||
vad_chunk_ms=max(50, int(os.getenv("FUNASR_VAD_CHUNK_MS", str(cls.vad_chunk_ms)))),
|
||||
chunk_size=chunk_size,
|
||||
encoder_chunk_look_back=max(
|
||||
0, int(os.getenv("FUNASR_ENCODER_LOOK_BACK", str(cls.encoder_chunk_look_back)))
|
||||
),
|
||||
decoder_chunk_look_back=max(
|
||||
0, int(os.getenv("FUNASR_DECODER_LOOK_BACK", str(cls.decoder_chunk_look_back)))
|
||||
),
|
||||
max_segment_sec=max(
|
||||
2.0, float(os.getenv("FUNASR_MAX_SEGMENT_SEC", str(cls.max_segment_sec)))
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FunASRSegment:
|
||||
"""A partial or final event returned by one browser session."""
|
||||
|
||||
text: str
|
||||
start_time_ms: float
|
||||
end_time_ms: float
|
||||
audio: bytes
|
||||
voiced_ms: float
|
||||
is_final: bool
|
||||
sentence_id: int
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
def _result_text(result: Any) -> str:
|
||||
"""Extract text from the list/dict result shapes used by FunASR."""
|
||||
if isinstance(result, list):
|
||||
return _result_text(result[0]) if result else ""
|
||||
if isinstance(result, dict):
|
||||
value = result.get("text")
|
||||
if value is None:
|
||||
value = result.get("value")
|
||||
return str(value or "").strip()
|
||||
return str(result or "").strip()
|
||||
|
||||
|
||||
def _vad_events(result: Any) -> list[tuple[float, float]]:
|
||||
"""Normalize FunASR streaming VAD output to start/end milliseconds."""
|
||||
if isinstance(result, list):
|
||||
result = result[0] if result else {}
|
||||
if isinstance(result, dict):
|
||||
result = result.get("value") or result.get("segments") or []
|
||||
if not isinstance(result, list):
|
||||
return []
|
||||
events: list[tuple[float, float]] = []
|
||||
for item in result:
|
||||
if not isinstance(item, (list, tuple)) or len(item) < 2:
|
||||
continue
|
||||
try:
|
||||
events.append((float(item[0]), float(item[1])))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
return events
|
||||
|
||||
|
||||
class FunASRModelService:
|
||||
"""Load FunASR once and create isolated streaming state per browser."""
|
||||
|
||||
native_partial_supported = True
|
||||
|
||||
def __init__(self, config: FunASRServiceConfig | None = None) -> None:
|
||||
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."""
|
||||
try:
|
||||
from funasr import AutoModel
|
||||
except ImportError as exc: # pragma: no cover - deployment-only branch
|
||||
raise RuntimeError(
|
||||
"FunASR is not installed; run pip install -r requirements.txt"
|
||||
) from exc
|
||||
|
||||
def load_models() -> tuple[Any, Any]:
|
||||
common = {"disable_pbar": True, "disable_log": True}
|
||||
asr = AutoModel(model=self.config.model, device=self.config.device, **common)
|
||||
vad = AutoModel(model=self.config.vad_model, device=self.config.vad_device, **common)
|
||||
return asr, vad
|
||||
|
||||
self.asr_model, self.vad_model = await asyncio.to_thread(load_models)
|
||||
LOGGER.info(
|
||||
"FunASR ready: model=%s vad=%s device=%s vad_device=%s",
|
||||
self.config.model,
|
||||
self.config.vad_model,
|
||||
self.config.device,
|
||||
self.config.vad_device,
|
||||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Release references so CUDA memory can be reclaimed on shutdown."""
|
||||
self.asr_model = None
|
||||
self.vad_model = None
|
||||
|
||||
def create_session(self) -> "FunASRRealtimeSession":
|
||||
"""Create a session with isolated VAD and ASR caches."""
|
||||
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."""
|
||||
if self.asr_model is None:
|
||||
raise RuntimeError("FunASR ASR model is not loaded")
|
||||
|
||||
def generate() -> str:
|
||||
return _result_text(self.asr_model.generate(input=audio, **status))
|
||||
|
||||
async with self.inference_lock:
|
||||
return await asyncio.to_thread(generate)
|
||||
|
||||
async def generate_vad(
|
||||
self,
|
||||
audio: Any,
|
||||
status: dict[str, Any],
|
||||
chunk_ms: int,
|
||||
) -> list[tuple[float, float]]:
|
||||
"""Run one streaming VAD chunk and normalize endpoint events."""
|
||||
if self.vad_model is None:
|
||||
raise RuntimeError("FunASR VAD model is not loaded")
|
||||
|
||||
def generate() -> list[tuple[float, float]]:
|
||||
result = self.vad_model.generate(input=audio, chunk_size=chunk_ms, **status)
|
||||
return _vad_events(result)
|
||||
|
||||
async with self.inference_lock:
|
||||
return await asyncio.to_thread(generate)
|
||||
|
||||
|
||||
class _StreamingASRTurn:
|
||||
"""One utterance using FunASR's ordered chunk/cache lifecycle."""
|
||||
|
||||
def __init__(self, service: FunASRModelService) -> None:
|
||||
self.service = service
|
||||
self.cache: dict[str, Any] = {}
|
||||
self.pending = bytearray()
|
||||
self.cumulative_text = ""
|
||||
self.last_chunk_text = ""
|
||||
self.pending_outputs: list[str] = []
|
||||
self.chunk_samples = max(1, service.config.chunk_size[1] * 960)
|
||||
|
||||
@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):
|
||||
return chunk
|
||||
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.
|
||||
while len(self.pending) >= chunk_bytes * 2:
|
||||
chunk = bytes(self.pending[:chunk_bytes])
|
||||
del self.pending[:chunk_bytes]
|
||||
text = await self.service.generate_asr(
|
||||
chunk,
|
||||
{
|
||||
"cache": self.cache,
|
||||
"is_final": False,
|
||||
"chunk_size": self.service.config.chunk_size,
|
||||
"encoder_chunk_look_back": self.service.config.encoder_chunk_look_back,
|
||||
"decoder_chunk_look_back": self.service.config.decoder_chunk_look_back,
|
||||
"batch_size": 1,
|
||||
},
|
||||
)
|
||||
self.cumulative_text = self._merge_chunk_text(
|
||||
self.cumulative_text, text, self.last_chunk_text
|
||||
)
|
||||
self.last_chunk_text = text
|
||||
if text:
|
||||
self.pending_outputs.append(self.cumulative_text)
|
||||
outputs = self.pending_outputs
|
||||
self.pending_outputs = []
|
||||
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()
|
||||
text = await self.service.generate_asr(
|
||||
chunk,
|
||||
{
|
||||
"cache": self.cache,
|
||||
"is_final": True,
|
||||
"chunk_size": self.service.config.chunk_size,
|
||||
"encoder_chunk_look_back": self.service.config.encoder_chunk_look_back,
|
||||
"decoder_chunk_look_back": self.service.config.decoder_chunk_look_back,
|
||||
"batch_size": 1,
|
||||
},
|
||||
)
|
||||
self.cumulative_text = self._merge_chunk_text(
|
||||
self.cumulative_text, text, self.last_chunk_text
|
||||
)
|
||||
self.last_chunk_text = text
|
||||
return self.cumulative_text.strip()
|
||||
|
||||
|
||||
class FunASRRealtimeSession:
|
||||
"""FunASR VAD + streaming ASR session used by one browser WebSocket."""
|
||||
|
||||
def __init__(self, service: FunASRModelService) -> None:
|
||||
self.service = service
|
||||
self.sample_rate = service.config.sample_rate
|
||||
self.vad_chunk_bytes = service.config.vad_chunk_ms * self.sample_rate * 2 // 1000
|
||||
self.vad_buffer = bytearray()
|
||||
self.vad_cache: dict[str, Any] = {}
|
||||
self.pre_roll = bytearray()
|
||||
self.pre_roll_max_bytes = 300 * self.sample_rate * 2 // 1000
|
||||
self.total_samples = 0
|
||||
self.speech_started = False
|
||||
self.segment_id = 0
|
||||
self.segment_start_ms = 0.0
|
||||
self.segment_audio = bytearray()
|
||||
self.segment_voiced_ms = 0.0
|
||||
self.asr_turn: _StreamingASRTurn | None = None
|
||||
|
||||
@staticmethod
|
||||
def _to_float32(pcm_bytes: bytes) -> Any:
|
||||
"""Convert browser PCM16 to the float waveform expected by 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,
|
||||
start_ms - len(self.pre_roll) / (self.sample_rate * 2) * 1000,
|
||||
)
|
||||
self.segment_audio = bytearray(self.pre_roll)
|
||||
self.segment_voiced_ms = 0.0
|
||||
self.asr_turn = _StreamingASRTurn(self.service)
|
||||
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."""
|
||||
self.total_samples += len(chunk) // 2
|
||||
events = await self.service.generate_vad(
|
||||
self._to_float32(chunk),
|
||||
{"cache": self.vad_cache, "is_final": False},
|
||||
self.service.config.vad_chunk_ms,
|
||||
)
|
||||
starts = [start for start, _ in events if start >= 0]
|
||||
ends = [end for _, end in events if end >= 0]
|
||||
results: list[FunASRSegment] = []
|
||||
if not self.speech_started and starts:
|
||||
self._start_segment(starts[0])
|
||||
|
||||
if self.speech_started:
|
||||
self.segment_audio.extend(chunk)
|
||||
self.segment_voiced_ms += len(chunk) / (self.sample_rate * 2) * 1000
|
||||
if self.asr_turn is not None:
|
||||
for text in await self.asr_turn.append(chunk):
|
||||
results.append(self._partial(text))
|
||||
if self._segment_duration_ms() >= self.service.config.max_segment_sec * 1000:
|
||||
results.append(await self._finish_segment("max_duration"))
|
||||
else:
|
||||
self.pre_roll.extend(chunk)
|
||||
del self.pre_roll[:-self.pre_roll_max_bytes]
|
||||
|
||||
if ends and self.speech_started:
|
||||
results.append(await self._finish_segment("vad_end"))
|
||||
return [event for event in results if event.text or event.is_final]
|
||||
|
||||
def _segment_duration_ms(self) -> float:
|
||||
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."""
|
||||
return FunASRSegment(
|
||||
text=text,
|
||||
start_time_ms=self.segment_start_ms,
|
||||
end_time_ms=self.segment_start_ms + self._segment_duration_ms(),
|
||||
audio=b"",
|
||||
voiced_ms=self.segment_voiced_ms,
|
||||
is_final=False,
|
||||
sentence_id=self.segment_id,
|
||||
)
|
||||
|
||||
async def _finish_segment(self, reason: str) -> FunASRSegment:
|
||||
"""Flush the ASR cache, then release completed segment audio."""
|
||||
text = await self.asr_turn.finish() if self.asr_turn is not None else ""
|
||||
result = FunASRSegment(
|
||||
text=text,
|
||||
start_time_ms=self.segment_start_ms,
|
||||
end_time_ms=self.segment_start_ms + self._segment_duration_ms(),
|
||||
audio=bytes(self.segment_audio),
|
||||
voiced_ms=self.segment_voiced_ms,
|
||||
is_final=True,
|
||||
sentence_id=self.segment_id,
|
||||
reason=reason,
|
||||
)
|
||||
self.segment_id += 1
|
||||
self.segment_audio.clear()
|
||||
self.asr_turn = None
|
||||
self.speech_started = False
|
||||
self.segment_voiced_ms = 0.0
|
||||
return result
|
||||
|
||||
async def feed(self, pcm_bytes: bytes) -> list[FunASRSegment]:
|
||||
"""Consume PCM16 and return FunASR partial/final events."""
|
||||
if len(pcm_bytes) % 2:
|
||||
raise ValueError("PCM16 音频必须包含完整的双字节采样")
|
||||
self.vad_buffer.extend(pcm_bytes)
|
||||
results: list[FunASRSegment] = []
|
||||
while len(self.vad_buffer) >= self.vad_chunk_bytes:
|
||||
chunk = bytes(self.vad_buffer[:self.vad_chunk_bytes])
|
||||
del self.vad_buffer[:self.vad_chunk_bytes]
|
||||
results.extend(await self._feed_vad_chunk(chunk))
|
||||
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)
|
||||
self.vad_buffer.clear()
|
||||
if not self.speech_started and chunk:
|
||||
self._start_segment(self.total_samples / self.sample_rate * 1000)
|
||||
self.total_samples += len(chunk) // 2
|
||||
if self.speech_started:
|
||||
self.segment_audio.extend(chunk)
|
||||
self.segment_voiced_ms += len(chunk) / (self.sample_rate * 2) * 1000
|
||||
if self.asr_turn is not None:
|
||||
for text in await self.asr_turn.append(chunk):
|
||||
results.append(self._partial(text))
|
||||
if self.speech_started:
|
||||
results.append(await self._finish_segment("eof"))
|
||||
self.vad_cache = {}
|
||||
self.pre_roll.clear()
|
||||
return [event for event in results if event.text or event.is_final]
|
||||
|
|
@ -1,957 +0,0 @@
|
|||
"""FunASR realtime WebSocket server adapted from the local FunASR checkout.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import websockets
|
||||
import time
|
||||
import numpy as np
|
||||
import argparse
|
||||
import ssl
|
||||
import os
|
||||
import wave
|
||||
import functools
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from scipy.spatial.distance import cosine
|
||||
|
||||
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."""
|
||||
raw_value = os.getenv(name, str(default)).strip()
|
||||
try:
|
||||
value = float(raw_value)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"{name} must be a number; got {raw_value!r}") from exc
|
||||
if not minimum <= value <= maximum:
|
||||
raise ValueError(f"{name} must be between {minimum} and {maximum}; got {value}")
|
||||
return value
|
||||
|
||||
|
||||
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)
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"{name} must be a positive integer; got {raw_value!r}") from exc
|
||||
if value <= 0:
|
||||
raise ValueError(f"{name} must be greater than zero; got {value}")
|
||||
return value
|
||||
|
||||
|
||||
# An explicit value disables FunASR's duration-based silence schedule for testing.
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
def to_python(obj):
|
||||
"""递归地把 numpy / torch 等类型转成纯 Python,可 JSON 序列化。"""
|
||||
try:
|
||||
import numpy as np # noqa
|
||||
import torch # noqa
|
||||
except Exception:
|
||||
np = None
|
||||
torch = None
|
||||
|
||||
if np is not None and isinstance(obj, np.generic):
|
||||
return obj.item()
|
||||
if np is not None and isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
if torch is not None and isinstance(obj, torch.Tensor):
|
||||
return obj.cpu().tolist()
|
||||
|
||||
if isinstance(obj, dict):
|
||||
return {k: to_python(v) for k, v in obj.items()}
|
||||
if isinstance(obj, (list, tuple)):
|
||||
return [to_python(v) for v in obj]
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--host", type=str, default="0.0.0.0", required=False, help="host ip")
|
||||
parser.add_argument("--port", type=int, default=10095, required=False, help="grpc server port")
|
||||
|
||||
parser.add_argument(
|
||||
"--asr_model",
|
||||
type=str,
|
||||
default="",
|
||||
help="Optional offline ASR model; empty means online-only mode.",
|
||||
)
|
||||
parser.add_argument("--asr_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--asr_model_online",
|
||||
type=str,
|
||||
default="iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--asr_model_online_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--vad_model",
|
||||
type=str,
|
||||
default="iic/speech_fsmn_vad_zh-cn-16k-common-pytorch",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--vad_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument(
|
||||
"--punc_model",
|
||||
type=str,
|
||||
default="",
|
||||
help="model from modelscope",
|
||||
)
|
||||
parser.add_argument("--punc_model_revision", type=str, default="v2.0.4", help="")
|
||||
|
||||
parser.add_argument("--ngpu", type=int, default=1, help="0 for cpu, 1 for gpu")
|
||||
parser.add_argument("--device", type=str, default="cuda", help="cuda, cpu")
|
||||
parser.add_argument("--vad_device", type=str, default=None, help="Optional VAD device override")
|
||||
parser.add_argument("--ncpu", type=int, default=4, help="cpu cores")
|
||||
parser.add_argument(
|
||||
"--enable_speaker_verification",
|
||||
action="store_true",
|
||||
help="Load native CAM++; disabled when CAM++ is hosted by the auxiliary service.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--certfile",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="certfile for ssl",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--keyfile",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="keyfile for ssl",
|
||||
)
|
||||
|
||||
# ====== 保存 2pass 离线阶段送入 ASR 的音频片段(排查 VAD 切分)======
|
||||
parser.add_argument(
|
||||
"--save_offline_segments",
|
||||
action="store_true",
|
||||
help="Save each offline (2pass) audio segment sent to offline ASR as wav for debugging VAD split.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save_offline_segments_dir",
|
||||
type=str,
|
||||
default="./offline_segments",
|
||||
help="Directory to save offline wav segments when --save_offline_segments is enabled.",
|
||||
)
|
||||
|
||||
# ====== 并发控制:核心新增 ======
|
||||
parser.add_argument(
|
||||
"--worker_threads",
|
||||
type=int,
|
||||
default=max(4, (os.cpu_count() or 4)),
|
||||
help="ThreadPoolExecutor max_workers. Used to offload blocking inference so event loop won't be blocked.",
|
||||
)
|
||||
parser.add_argument("--concurrent_vad", type=int, default=4, help="Max concurrent VAD generate() calls.")
|
||||
parser.add_argument("--concurrent_asr_online", type=int, default=4, help="Max concurrent streaming ASR generate() calls.")
|
||||
parser.add_argument("--concurrent_asr_offline", type=int, default=2, help="Max concurrent offline ASR generate() calls.")
|
||||
parser.add_argument("--concurrent_punc", type=int, default=1, help="Max concurrent punctuation generate() calls.")
|
||||
parser.add_argument("--concurrent_sv", type=int, default=1, help="Max concurrent speaker verification generate() calls.")
|
||||
parser.add_argument(
|
||||
"--speaker_db_reload_sec",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Reload speaker_db.json at most once every N seconds (avoid frequent disk IO).",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
websocket_users = set()
|
||||
SPEAKER_DB_PATH = os.path.join(os.path.dirname(__file__), "speaker_db.json")
|
||||
|
||||
|
||||
def _ensure_dir(p: str):
|
||||
try:
|
||||
os.makedirs(p, exist_ok=True)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _pcm_duration_ms(pcm_bytes: bytes, fs: int, ch: int = 1, sampwidth: int = 2) -> int:
|
||||
"""根据 fs/ch/sampwidth 计算 PCM 时长,避免写死 16k -> 32 bytes/ms。"""
|
||||
if not pcm_bytes:
|
||||
return 0
|
||||
bytes_per_ms = (fs * ch * sampwidth) / 1000.0
|
||||
if bytes_per_ms <= 0:
|
||||
return 0
|
||||
return int(len(pcm_bytes) / bytes_per_ms)
|
||||
|
||||
|
||||
def _safe_int(v, default):
|
||||
try:
|
||||
return int(v)
|
||||
except Exception:
|
||||
return default
|
||||
|
||||
|
||||
# ========= speaker db:加缓存,避免每段都读盘 =========
|
||||
_SPEAKER_DB_CACHE = {}
|
||||
_SPEAKER_DB_CACHE_TS = 0.0
|
||||
|
||||
|
||||
def _load_speaker_db_sync():
|
||||
if not os.path.exists(SPEAKER_DB_PATH):
|
||||
return {}
|
||||
try:
|
||||
with open(SPEAKER_DB_PATH, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def get_speaker_db_cached(now_ts: float, reload_sec: int):
|
||||
global _SPEAKER_DB_CACHE, _SPEAKER_DB_CACHE_TS
|
||||
if (now_ts - _SPEAKER_DB_CACHE_TS) >= max(1, int(reload_sec)):
|
||||
_SPEAKER_DB_CACHE = _load_speaker_db_sync()
|
||||
_SPEAKER_DB_CACHE_TS = now_ts
|
||||
return _SPEAKER_DB_CACHE or {}
|
||||
|
||||
|
||||
def _save_wav_sync(out_path: str, audio_bytes: bytes, fs: int, ch: int, sampwidth: int):
|
||||
with wave.open(out_path, "wb") as wf:
|
||||
wf.setnchannels(ch)
|
||||
wf.setsampwidth(sampwidth)
|
||||
wf.setframerate(fs)
|
||||
wf.writeframes(audio_bytes)
|
||||
|
||||
|
||||
def save_offline_wav_segment_sync(websocket, audio_bytes: bytes, reason: str = "offline"):
|
||||
"""
|
||||
保存离线阶段送入 ASR 的音频片段,方便人工试听排查 VAD 切分是否正确。
|
||||
约定:audio_bytes 为 单声道 PCM16 little-endian(默认 16k)。
|
||||
(注意:这是同步函数,外层会放线程池执行)
|
||||
"""
|
||||
if not getattr(websocket, "save_offline_segments", False):
|
||||
return
|
||||
if "2pass" not in (getattr(websocket, "mode", "") or ""):
|
||||
return
|
||||
if not audio_bytes:
|
||||
return
|
||||
|
||||
fs = int(getattr(websocket, "audio_fs", 16000) or 16000)
|
||||
ch = 1
|
||||
sampwidth = 2 # int16
|
||||
|
||||
# int16 对齐
|
||||
if len(audio_bytes) % 2 == 1:
|
||||
audio_bytes = audio_bytes[:-1]
|
||||
if not audio_bytes:
|
||||
return
|
||||
|
||||
seg_idx = int(getattr(websocket, "offline_seg_idx", 0))
|
||||
websocket.offline_seg_idx = seg_idx + 1
|
||||
|
||||
duration_ms = _pcm_duration_ms(audio_bytes, fs=fs, ch=ch, sampwidth=sampwidth)
|
||||
|
||||
base_dir = getattr(websocket, "offline_save_dir", args.save_offline_segments_dir)
|
||||
_ensure_dir(base_dir)
|
||||
|
||||
wav_name = (getattr(websocket, "wav_name", "microphone") or "microphone").replace("/", "_")
|
||||
ts = int(time.time() * 1000)
|
||||
fname = f"{wav_name}_{ts}_seg{seg_idx:04d}_{reason}_{duration_ms}ms.wav"
|
||||
out_path = os.path.join(base_dir, fname)
|
||||
|
||||
try:
|
||||
_save_wav_sync(out_path, audio_bytes, fs=fs, ch=ch, sampwidth=sampwidth)
|
||||
print(f"[SAVE_OFFLINE_SEG] {out_path} ({duration_ms} ms, {len(audio_bytes)} bytes)")
|
||||
except Exception as e:
|
||||
print(f"[SAVE_OFFLINE_SEG] failed: {e}")
|
||||
|
||||
|
||||
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,
|
||||
model_revision=args.asr_model_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
if args.asr_model
|
||||
else None
|
||||
)
|
||||
|
||||
# streaming asr
|
||||
model_asr_streaming = AutoModel(
|
||||
model=args.asr_model_online,
|
||||
model_revision=args.asr_model_online_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
|
||||
# vad
|
||||
model_vad = AutoModel(
|
||||
model=args.vad_model,
|
||||
model_revision=args.vad_model_revision,
|
||||
ngpu=args.ngpu if (args.vad_device or args.device).startswith("cuda") else 0,
|
||||
ncpu=args.ncpu,
|
||||
device=args.vad_device or args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
|
||||
# punc
|
||||
if args.punc_model != "":
|
||||
model_punc = AutoModel(
|
||||
model=args.punc_model,
|
||||
model_revision=args.punc_model_revision,
|
||||
ngpu=args.ngpu,
|
||||
ncpu=args.ncpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
else:
|
||||
model_punc = None
|
||||
|
||||
# CAM++ is loaded by the auxiliary service, avoiding a second GPU copy here.
|
||||
model_sv = (
|
||||
AutoModel(
|
||||
model="iic/speech_campplus_sv_zh-cn_16k-common",
|
||||
ngpu=args.ngpu,
|
||||
device=args.device,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
)
|
||||
if args.enable_speaker_verification
|
||||
else None
|
||||
)
|
||||
|
||||
print("model loaded! (now supports multi-client with non-blocking inference)")
|
||||
|
||||
|
||||
# ====== 线程池 + 并发阈值(核心)======
|
||||
EXECUTOR = ThreadPoolExecutor(max_workers=int(args.worker_threads))
|
||||
|
||||
SEM_VAD = asyncio.Semaphore(max(1, int(args.concurrent_vad)))
|
||||
SEM_ASR_ONLINE = asyncio.Semaphore(max(1, int(args.concurrent_asr_online)))
|
||||
SEM_ASR_OFFLINE = asyncio.Semaphore(max(1, int(args.concurrent_asr_offline)))
|
||||
SEM_PUNC = asyncio.Semaphore(max(1, int(args.concurrent_punc)))
|
||||
SEM_SV = asyncio.Semaphore(max(1, int(args.concurrent_sv)))
|
||||
SEM_WAV = asyncio.Semaphore(max(1, 4)) # 保存 wav 一般不需要太大
|
||||
|
||||
|
||||
async def run_blocking(fn, *a, sem: asyncio.Semaphore | None = None, **kw):
|
||||
"""
|
||||
把阻塞函数丢线程池执行,避免卡 event loop。
|
||||
sem 用于限流(避免 GPU / 模型被打爆)。
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
call = functools.partial(fn, *a, **kw)
|
||||
if sem is None:
|
||||
return await loop.run_in_executor(EXECUTOR, call)
|
||||
async with sem:
|
||||
return await loop.run_in_executor(EXECUTOR, call)
|
||||
|
||||
|
||||
def _generate_sync(model, audio_or_text, status_dict):
|
||||
# 注意:status_dict 里包含 cache,会被 generate 更新
|
||||
return model.generate(input=audio_or_text, **status_dict)
|
||||
|
||||
|
||||
async def ws_reset(websocket):
|
||||
print("ws reset now, total num is ", len(websocket_users))
|
||||
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = True
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
websocket.status_dict_vad["is_final"] = True
|
||||
websocket.status_dict_punc["cache"] = {}
|
||||
|
||||
await websocket.close()
|
||||
|
||||
|
||||
async def clear_websocket():
|
||||
for websocket in list(websocket_users):
|
||||
await ws_reset(websocket)
|
||||
websocket_users.clear()
|
||||
|
||||
|
||||
async def ws_serve(websocket, path=None):
|
||||
# websockets 新版本不会传 path,这里做兼容
|
||||
if path is None:
|
||||
path = getattr(websocket, "path", None)
|
||||
frames = []
|
||||
frames_asr = []
|
||||
frames_asr_online = []
|
||||
pending_offline_audio = []
|
||||
global websocket_users
|
||||
websocket_users.add(websocket)
|
||||
|
||||
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.
|
||||
websocket.status_dict_vad = {
|
||||
"cache": {},
|
||||
"is_final": False,
|
||||
"max_end_silence_time": VAD_MAX_END_SILENCE_MS,
|
||||
"speech_noise_thres": VAD_SPEECH_NOISE_THRESHOLD,
|
||||
}
|
||||
websocket.status_dict_punc = {"cache": {}}
|
||||
|
||||
websocket.chunk_interval = 10
|
||||
websocket.sentence_strategy = 0
|
||||
websocket.vad_pre_idx = 0
|
||||
speech_start = False
|
||||
speech_end_i = -1
|
||||
online_needs_finalization = False
|
||||
session_errors = []
|
||||
|
||||
websocket.wav_name = "microphone"
|
||||
websocket.mode = "2pass"
|
||||
websocket.is_speaking = True # ✅ 默认初始化,避免 AttributeError
|
||||
|
||||
# 保存离线片段
|
||||
websocket.audio_fs = 16000
|
||||
websocket.offline_seg_idx = 0
|
||||
websocket.save_offline_segments = bool(args.save_offline_segments)
|
||||
websocket.offline_save_dir = args.save_offline_segments_dir
|
||||
if websocket.save_offline_segments:
|
||||
_ensure_dir(websocket.offline_save_dir)
|
||||
print(f"[SAVE_OFFLINE_SEG] enabled, dir={websocket.offline_save_dir}")
|
||||
|
||||
print("new user connected", flush=True)
|
||||
|
||||
def record_error(message):
|
||||
if message not in session_errors:
|
||||
session_errors.append(message)
|
||||
|
||||
async def finalize_online_segment():
|
||||
nonlocal frames_asr_online, online_needs_finalization
|
||||
|
||||
if websocket.mode not in ("2pass", "online") or not online_needs_finalization:
|
||||
return
|
||||
|
||||
websocket.status_dict_asr_online["is_final"] = True
|
||||
try:
|
||||
await async_asr_online(websocket, b"".join(frames_asr_online))
|
||||
except Exception as e:
|
||||
print("error in final asr streaming:", e)
|
||||
record_error(f"online inference failed: {e}")
|
||||
|
||||
frames_asr_online = []
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = False
|
||||
online_needs_finalization = False
|
||||
|
||||
async def finish_input(send_end_ack):
|
||||
nonlocal frames, frames_asr, frames_asr_online, pending_offline_audio
|
||||
nonlocal speech_start, speech_end_i, online_needs_finalization
|
||||
|
||||
await finalize_online_segment()
|
||||
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
audio_in = b"".join(frames_asr)
|
||||
if not audio_in:
|
||||
audio_in = b"".join(pending_offline_audio)
|
||||
|
||||
if audio_in:
|
||||
if websocket.save_offline_segments and audio_in:
|
||||
try:
|
||||
await run_blocking(
|
||||
save_offline_wav_segment_sync,
|
||||
websocket,
|
||||
audio_in,
|
||||
"not_speaking",
|
||||
sem=SEM_WAV,
|
||||
)
|
||||
except Exception as e:
|
||||
print("[SAVE_OFFLINE_SEG] async failed:", e)
|
||||
|
||||
try:
|
||||
await async_asr(websocket, audio_in)
|
||||
pending_offline_audio = []
|
||||
except Exception as e:
|
||||
print("error in final asr offline:", e)
|
||||
record_error(f"offline inference failed: {e}")
|
||||
|
||||
errors = list(session_errors)
|
||||
|
||||
frames = []
|
||||
frames_asr = []
|
||||
frames_asr_online = []
|
||||
pending_offline_audio = []
|
||||
speech_start = False
|
||||
speech_end_i = -1
|
||||
online_needs_finalization = False
|
||||
websocket.vad_pre_idx = 0
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
|
||||
if send_end_ack:
|
||||
acknowledgement = {
|
||||
"mode": websocket.mode,
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": not errors,
|
||||
"is_end": True,
|
||||
}
|
||||
if errors:
|
||||
acknowledgement["error"] = "; ".join(errors)
|
||||
await websocket.send(
|
||||
json.dumps(acknowledgement, ensure_ascii=False)
|
||||
)
|
||||
session_errors.clear()
|
||||
elif errors:
|
||||
raise RuntimeError("; ".join(errors))
|
||||
|
||||
try:
|
||||
async for message in websocket:
|
||||
# ========== 1) 先处理“文本配置消息” ==========
|
||||
if isinstance(message, str):
|
||||
try:
|
||||
messagejson = json.loads(message)
|
||||
except Exception as e:
|
||||
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:
|
||||
websocket.is_speaking = bool(messagejson["is_speaking"])
|
||||
websocket.status_dict_asr_online["is_final"] = (not websocket.is_speaking)
|
||||
end_of_input = not websocket.is_speaking
|
||||
|
||||
if "chunk_interval" in messagejson:
|
||||
websocket.chunk_interval = _safe_int(
|
||||
messagejson["chunk_interval"], websocket.chunk_interval
|
||||
)
|
||||
|
||||
if "sentence_strategy" in messagejson:
|
||||
# Map the Tencent selector to FunASR's VAD endpoint duration.
|
||||
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"] = (
|
||||
VAD_PARAGRAPH_MAX_END_SILENCE_MS
|
||||
if websocket.sentence_strategy == 1
|
||||
else VAD_MAX_END_SILENCE_MS
|
||||
)
|
||||
|
||||
if "wav_name" in messagejson:
|
||||
websocket.wav_name = messagejson.get("wav_name") or websocket.wav_name
|
||||
|
||||
if "chunk_size" in messagejson:
|
||||
chunk_size = messagejson["chunk_size"]
|
||||
if isinstance(chunk_size, str):
|
||||
chunk_size = [x.strip() for x in chunk_size.split(",") if x.strip()]
|
||||
websocket.status_dict_asr_online["chunk_size"] = [int(x) for x in chunk_size]
|
||||
|
||||
if "encoder_chunk_look_back" in messagejson:
|
||||
websocket.status_dict_asr_online["encoder_chunk_look_back"] = messagejson[
|
||||
"encoder_chunk_look_back"
|
||||
]
|
||||
|
||||
if "decoder_chunk_look_back" in messagejson:
|
||||
websocket.status_dict_asr_online["decoder_chunk_look_back"] = messagejson[
|
||||
"decoder_chunk_look_back"
|
||||
]
|
||||
|
||||
if "hotwords" in messagejson:
|
||||
hotword_data = messagejson["hotwords"]
|
||||
websocket.status_dict_asr["hotword"] = hotword_data
|
||||
websocket.status_dict_asr_online["hotword"] = hotword_data
|
||||
print(f"热词已更新: {hotword_data}")
|
||||
|
||||
if "mode" in messagejson:
|
||||
requested_mode = messagejson["mode"]
|
||||
if requested_mode and requested_mode not in ("online", "offline", "2pass"):
|
||||
websocket.mode = requested_mode
|
||||
record_error(f"unsupported mode: {requested_mode!r}")
|
||||
else:
|
||||
websocket.mode = requested_mode or websocket.mode
|
||||
|
||||
if "audio_fs" in messagejson:
|
||||
websocket.audio_fs = _safe_int(messagejson["audio_fs"], 16000)
|
||||
|
||||
if end_of_input:
|
||||
await finish_input(send_end_ack=bool(messagejson.get("is_end")))
|
||||
|
||||
continue
|
||||
|
||||
# ========== 2) 处理“二进制音频消息” ==========
|
||||
if websocket.mode not in ("online", "offline", "2pass"):
|
||||
continue
|
||||
|
||||
if "chunk_size" not in websocket.status_dict_asr_online:
|
||||
print("[WARN] chunk_size not set yet, skip audio frame (send config first).")
|
||||
record_error("audio frame discarded: chunk_size is not configured")
|
||||
continue
|
||||
|
||||
try:
|
||||
websocket.status_dict_vad["chunk_size"] = int(
|
||||
websocket.status_dict_asr_online["chunk_size"][1] * 60 / websocket.chunk_interval
|
||||
)
|
||||
except Exception as e:
|
||||
print("[WARN] set vad chunk_size failed:", e)
|
||||
record_error(f"audio frame discarded: invalid VAD chunk_size: {e}")
|
||||
continue
|
||||
|
||||
pcm = message
|
||||
frames.append(pcm)
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
pending_offline_audio.append(pcm)
|
||||
|
||||
duration_ms = _pcm_duration_ms(pcm, fs=websocket.audio_fs, ch=1, sampwidth=2)
|
||||
websocket.vad_pre_idx += duration_ms
|
||||
|
||||
# online asr
|
||||
frames_asr_online.append(pcm)
|
||||
if websocket.mode in ("2pass", "online"):
|
||||
online_needs_finalization = True
|
||||
websocket.status_dict_asr_online["is_final"] = (speech_end_i != -1)
|
||||
|
||||
if (len(frames_asr_online) % websocket.chunk_interval == 0) or websocket.status_dict_asr_online["is_final"]:
|
||||
if websocket.mode in ("2pass", "online"):
|
||||
audio_in = b"".join(frames_asr_online)
|
||||
try:
|
||||
await async_asr_online(websocket, audio_in)
|
||||
except Exception as e:
|
||||
print(f"error in asr streaming, {websocket.status_dict_asr_online}")
|
||||
record_error(f"online inference failed: {e}")
|
||||
frames_asr_online = []
|
||||
|
||||
if speech_start:
|
||||
frames_asr.append(pcm)
|
||||
|
||||
# vad online
|
||||
try:
|
||||
speech_start_i, speech_end_i = await async_vad(websocket, pcm)
|
||||
except Exception as e:
|
||||
print("error in vad:", e)
|
||||
record_error(f"vad inference failed: {e}")
|
||||
speech_start_i, speech_end_i = -1, -1
|
||||
|
||||
# 把 FunASR VAD 的绝对音频时间发给桥接层,用于去掉首尾静音,
|
||||
# 让 CAM++ 滑窗只分析当前 VAD turn,而不是整段会话的缓冲音频。
|
||||
if speech_start_i != -1 or speech_end_i != -1:
|
||||
await websocket.send(
|
||||
json.dumps(
|
||||
{
|
||||
"event": "vad",
|
||||
"speech_start_ms": speech_start_i if speech_start_i != -1 else None,
|
||||
"speech_end_ms": speech_end_i if speech_end_i != -1 else None,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
)
|
||||
|
||||
if speech_start_i != -1:
|
||||
speech_start = True
|
||||
if duration_ms > 0:
|
||||
beg_bias = (websocket.vad_pre_idx - speech_start_i) // duration_ms
|
||||
else:
|
||||
beg_bias = 0
|
||||
frames_pre = frames[-beg_bias:] if beg_bias > 0 else []
|
||||
frames_asr = []
|
||||
frames_asr.extend(frames_pre)
|
||||
|
||||
# ========== 3) 2pass:离线阶段触发点 ==========
|
||||
if (speech_end_i != -1) or (not websocket.is_speaking):
|
||||
await finalize_online_segment()
|
||||
|
||||
if websocket.mode in ("2pass", "offline"):
|
||||
audio_in = b"".join(frames_asr)
|
||||
if not audio_in and speech_end_i != -1:
|
||||
audio_in = b"".join(pending_offline_audio)
|
||||
reason = "vad_end" if speech_end_i != -1 else "not_speaking"
|
||||
|
||||
# 保存 wav:放线程池,避免磁盘 IO 卡 loop
|
||||
if websocket.save_offline_segments and audio_in:
|
||||
try:
|
||||
await run_blocking(
|
||||
save_offline_wav_segment_sync,
|
||||
websocket,
|
||||
audio_in,
|
||||
reason,
|
||||
sem=SEM_WAV,
|
||||
)
|
||||
except Exception as e:
|
||||
print("[SAVE_OFFLINE_SEG] async failed:", e)
|
||||
|
||||
if audio_in:
|
||||
try:
|
||||
await async_asr(websocket, audio_in)
|
||||
pending_offline_audio = []
|
||||
except Exception as e:
|
||||
print("error in asr offline:", e)
|
||||
record_error(f"offline inference failed: {e}")
|
||||
|
||||
frames_asr = []
|
||||
speech_start = False
|
||||
frames_asr_online = []
|
||||
websocket.status_dict_asr_online["cache"] = {}
|
||||
websocket.status_dict_asr_online["is_final"] = False
|
||||
online_needs_finalization = False
|
||||
speech_end_i = -1
|
||||
|
||||
if not websocket.is_speaking:
|
||||
websocket.vad_pre_idx = 0
|
||||
frames = []
|
||||
websocket.status_dict_vad["cache"] = {}
|
||||
else:
|
||||
frames = frames[-20:]
|
||||
|
||||
except websockets.ConnectionClosed:
|
||||
print("ConnectionClosed...", websocket_users, flush=True)
|
||||
await ws_reset(websocket)
|
||||
if websocket in websocket_users:
|
||||
websocket_users.remove(websocket)
|
||||
except websockets.InvalidState:
|
||||
print("InvalidState...")
|
||||
try:
|
||||
await ws_reset(websocket)
|
||||
except Exception:
|
||||
pass
|
||||
websocket_users.discard(websocket)
|
||||
except Exception as e:
|
||||
print("Exception:", e)
|
||||
try:
|
||||
await ws_reset(websocket)
|
||||
except Exception:
|
||||
pass
|
||||
if websocket in websocket_users:
|
||||
websocket_users.remove(websocket)
|
||||
|
||||
|
||||
# ===================== 推理:全部改为“线程池 + 限流” =====================
|
||||
|
||||
async def async_vad(websocket, audio_in: bytes):
|
||||
# model_vad.generate 是阻塞的,必须 offload
|
||||
out = await run_blocking(_generate_sync, model_vad, audio_in, websocket.status_dict_vad, sem=SEM_VAD)
|
||||
segments_result = out[0].get("value", [])
|
||||
|
||||
speech_start = -1
|
||||
speech_end = -1
|
||||
|
||||
if len(segments_result) == 0 or len(segments_result) > 1:
|
||||
return speech_start, speech_end
|
||||
if segments_result[0][0] != -1:
|
||||
speech_start = segments_result[0][0]
|
||||
if segments_result[0][1] != -1:
|
||||
speech_end = segments_result[0][1]
|
||||
return speech_start, speech_end
|
||||
|
||||
|
||||
def _sv_and_match_sync(audio_in: bytes, reload_sec: int):
|
||||
"""
|
||||
同步执行:SV embedding + speaker_db 匹配
|
||||
返回 (spk_name, best_score)
|
||||
"""
|
||||
spk_name = "unknown"
|
||||
best_score = 0.0
|
||||
|
||||
sv_out = model_sv.generate(input=audio_in, embedding=True)[0]
|
||||
embedding = sv_out["spk_embedding"][0].cpu().numpy()
|
||||
|
||||
now_ts = time.time()
|
||||
local_speaker_db = get_speaker_db_cached(now_ts, reload_sec=reload_sec)
|
||||
if local_speaker_db:
|
||||
for name, ref_embedding in local_speaker_db.items():
|
||||
if ref_embedding is None:
|
||||
continue
|
||||
arr = np.array(ref_embedding, dtype=np.float32)
|
||||
similarity = 1.0 - cosine(embedding, arr)
|
||||
print("sv similarity with {}: {}".format(name, similarity))
|
||||
if similarity > best_score and similarity > 0.2:
|
||||
best_score = similarity
|
||||
spk_name = name
|
||||
|
||||
return spk_name, float(best_score)
|
||||
|
||||
|
||||
async def async_asr(websocket, audio_in: bytes):
|
||||
mode = "2pass-offline" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
if model_asr is None:
|
||||
raise RuntimeError("offline ASR is disabled; use FunASR online mode")
|
||||
|
||||
if len(audio_in) <= 0:
|
||||
message = {
|
||||
"mode": mode,
|
||||
"text": "",
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
return
|
||||
|
||||
# 1) ASR(阻塞,线程池执行)
|
||||
rec_result_list = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr,
|
||||
audio_in,
|
||||
websocket.status_dict_asr,
|
||||
sem=SEM_ASR_OFFLINE,
|
||||
)
|
||||
rec_result = rec_result_list[0]
|
||||
|
||||
print("offline_asr, raw:", rec_result)
|
||||
print("offline_asr, keys:", rec_result.keys())
|
||||
|
||||
text = rec_result.get("text", "")
|
||||
timestamp = rec_result.get("timestamp", None)
|
||||
sentence_info = rec_result.get("sentence_info", None)
|
||||
|
||||
# 2) 声纹识别(阻塞,线程池执行)
|
||||
spk_name = "unknown"
|
||||
best_score = 0.0
|
||||
try:
|
||||
spk_name, best_score = await run_blocking(
|
||||
_sv_and_match_sync,
|
||||
audio_in,
|
||||
int(args.speaker_db_reload_sec),
|
||||
sem=SEM_SV,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"声纹识别失败: {e}")
|
||||
|
||||
# 3) 标点(阻塞,线程池执行)
|
||||
punc_array = None
|
||||
if model_punc is not None and len(text) > 0:
|
||||
try:
|
||||
# punc 只对文本处理
|
||||
punc_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_punc,
|
||||
text,
|
||||
websocket.status_dict_punc,
|
||||
sem=SEM_PUNC,
|
||||
)
|
||||
punc_result = punc_out[0]
|
||||
print("offline, after punc", punc_result)
|
||||
|
||||
if "text" in punc_result and punc_result["text"]:
|
||||
text = punc_result["text"]
|
||||
if "punc_array" in punc_result:
|
||||
punc_array = punc_result["punc_array"]
|
||||
except Exception as e:
|
||||
print("punc failed:", e)
|
||||
|
||||
# 4) 构造最终 message
|
||||
if len(text) > 0:
|
||||
print("======offline final text:", text)
|
||||
message = {
|
||||
"mode": mode,
|
||||
"spk_name": spk_name,
|
||||
"spk_score": float(best_score),
|
||||
"text": text,
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
if timestamp is not None:
|
||||
message["timestamp"] = to_python(timestamp)
|
||||
if sentence_info is not None:
|
||||
message["sentence_info"] = to_python(sentence_info)
|
||||
if punc_array is not None:
|
||||
message["punc_array"] = to_python(punc_array)
|
||||
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
else:
|
||||
message = {
|
||||
"mode": mode,
|
||||
"spk_name": spk_name,
|
||||
"spk_score": float(best_score),
|
||||
"text": "",
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": True,
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
|
||||
|
||||
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 也是阻塞:线程池执行
|
||||
rec_out = await run_blocking(
|
||||
_generate_sync,
|
||||
model_asr_streaming,
|
||||
audio_in,
|
||||
websocket.status_dict_asr_online,
|
||||
sem=SEM_ASR_ONLINE,
|
||||
)
|
||||
rec_result = rec_out[0]
|
||||
|
||||
# 2pass:online 只要 partial,不发 final(final 交给 offline)
|
||||
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 状态。
|
||||
if rec_result.get("text") or is_final:
|
||||
mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
message = {
|
||||
"mode": mode,
|
||||
"text": rec_result.get("text", ""),
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": is_final,
|
||||
}
|
||||
await websocket.send(json.dumps(message, ensure_ascii=False))
|
||||
|
||||
|
||||
# ===================== 启动服务 =====================
|
||||
|
||||
async def main():
|
||||
if len(args.certfile) > 0:
|
||||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
ssl_context.load_cert_chain(args.certfile, keyfile=args.keyfile)
|
||||
server = await websockets.serve(
|
||||
ws_serve,
|
||||
args.host,
|
||||
args.port,
|
||||
subprotocols=["binary"],
|
||||
ping_interval=None,
|
||||
ssl=ssl_context,
|
||||
)
|
||||
else:
|
||||
server = await websockets.serve(
|
||||
ws_serve,
|
||||
args.host,
|
||||
args.port,
|
||||
subprotocols=["binary"],
|
||||
ping_interval=None,
|
||||
)
|
||||
|
||||
print(f"WS server started at ws(s)://{args.host}:{args.port}")
|
||||
await server.wait_closed()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
asyncio.run(main())
|
||||
finally:
|
||||
try:
|
||||
EXECUTOR.shutdown(wait=False, cancel_futures=True)
|
||||
except Exception:
|
||||
pass
|
||||
|
|
@ -1,715 +0,0 @@
|
|||
"""Translate the unchanged Tencent demo protocol to FunASR's native online WS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from aiohttp import WSMsgType, web
|
||||
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.
|
||||
from websockets import connect as websocket_connect
|
||||
|
||||
try:
|
||||
from .auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||||
except ImportError:
|
||||
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
|
||||
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
WEB_HOST = os.getenv("WEB_HOST", "0.0.0.0")
|
||||
WEB_PORT = int(os.getenv("WEB_PORT", "8082"))
|
||||
NATIVE_WS_URL = os.getenv("FUNASR_NATIVE_WS_URL", "ws://127.0.0.1:10095")
|
||||
CHUNK_SIZE = tuple(
|
||||
int(part.strip()) for part in os.getenv("FUNASR_CHUNK_SIZE", "0,10,5").split(",")
|
||||
)
|
||||
if len(CHUNK_SIZE) != 3:
|
||||
CHUNK_SIZE = (0, 10, 5)
|
||||
CHUNK_INTERVAL = max(1, int(os.getenv("FUNASR_CHUNK_INTERVAL", "10")))
|
||||
SAMPLE_RATE = 16000
|
||||
PCM_BYTES_PER_MS = SAMPLE_RATE * 2 / 1000
|
||||
FRAME_BYTES = max(2, round(60 * CHUNK_SIZE[1] / CHUNK_INTERVAL * PCM_BYTES_PER_MS))
|
||||
MAX_SPEAKER_AUDIO_BYTES = 60 * SAMPLE_RATE * 2
|
||||
MIN_SPEAKER_AUDIO_BYTES = int(0.8 * SAMPLE_RATE * 2)
|
||||
FINALIZE_TIMEOUT_SECONDS = max(30, int(os.getenv("FUNASR_FINALIZE_TIMEOUT_SECONDS", "300")))
|
||||
TURN_SENTENCE_ID_STRIDE = 100
|
||||
|
||||
AUXILIARY_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
|
||||
SESSION_REGISTRY_KEY = web.AppKey("sessions", dict)
|
||||
|
||||
|
||||
class IncrementalWavDecoder:
|
||||
"""Read a streamed PCM WAV header and yield its 16 kHz mono PCM payload."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.buffer = bytearray()
|
||||
self.header_done = False
|
||||
self.format_valid = False
|
||||
self.data_remaining: int | None = None
|
||||
|
||||
def feed(self, data: bytes) -> bytes:
|
||||
if self.header_done:
|
||||
if self.data_remaining is None:
|
||||
return data
|
||||
payload = data[: self.data_remaining]
|
||||
self.data_remaining -= len(payload)
|
||||
return payload
|
||||
|
||||
self.buffer.extend(data)
|
||||
if len(self.buffer) < 12:
|
||||
return b""
|
||||
if self.buffer[:4] != b"RIFF" or self.buffer[8:12] != b"WAVE":
|
||||
raise ValueError("WAV file must use a RIFF/WAVE container")
|
||||
del self.buffer[:12]
|
||||
|
||||
while len(self.buffer) >= 8:
|
||||
kind = bytes(self.buffer[:4])
|
||||
size = int.from_bytes(self.buffer[4:8], "little")
|
||||
if size > 1024 * 1024:
|
||||
raise ValueError("WAV header chunk is unexpectedly large")
|
||||
full_size = 8 + size + (size % 2)
|
||||
if len(self.buffer) < full_size:
|
||||
return b""
|
||||
body = bytes(self.buffer[8 : 8 + size])
|
||||
if kind == b"fmt ":
|
||||
if size < 16:
|
||||
raise ValueError("WAV fmt chunk is incomplete")
|
||||
fmt = (
|
||||
int.from_bytes(body[0:2], "little"),
|
||||
int.from_bytes(body[2:4], "little"),
|
||||
int.from_bytes(body[4:8], "little"),
|
||||
int.from_bytes(body[14:16], "little"),
|
||||
)
|
||||
if fmt != (1, 1, SAMPLE_RATE, 16):
|
||||
raise ValueError("WAV must be PCM16, mono, 16 kHz")
|
||||
self.format_valid = True
|
||||
if kind == b"data":
|
||||
if not self.format_valid or size % 2:
|
||||
raise ValueError("WAV must be PCM16, mono, 16 kHz")
|
||||
self.header_done = True
|
||||
self.data_remaining = size
|
||||
del self.buffer[:8]
|
||||
payload = bytes(self.buffer[:size])
|
||||
del self.buffer[: min(size, len(self.buffer))]
|
||||
self.data_remaining -= len(payload)
|
||||
return payload
|
||||
del self.buffer[:full_size]
|
||||
return b""
|
||||
|
||||
def finish(self) -> None:
|
||||
if not self.header_done or self.data_remaining not in (None, 0):
|
||||
raise ValueError("WAV ended before its complete PCM data chunk arrived")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SpeakerJob:
|
||||
sentence_id: int
|
||||
text: str
|
||||
audio: bytes
|
||||
start_time_ms: float
|
||||
end_time_ms: float
|
||||
|
||||
|
||||
def split_text_by_speaker_segments(
|
||||
text: str,
|
||||
segments: list[dict[str, Any]],
|
||||
turn_start_ms: float,
|
||||
turn_end_ms: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""按 CAM++ 时间段近似拆分 ASR 文本,并保留每段稳定说话人标签。"""
|
||||
usable: list[dict[str, Any]] = []
|
||||
for segment in sorted(
|
||||
segments,
|
||||
key=lambda item: float(item.get("start_time", turn_start_ms)),
|
||||
):
|
||||
try:
|
||||
start = max(turn_start_ms, float(segment.get("start_time", turn_start_ms)))
|
||||
end = min(turn_end_ms, float(segment.get("end_time", turn_end_ms)))
|
||||
speaker_id = int(segment.get("speaker_id", -1))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if end <= start or speaker_id < 0:
|
||||
continue
|
||||
speaker = {
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_name": str(segment.get("speaker_name") or f"说话人 {speaker_id + 1}"),
|
||||
"speaker_confidence": float(segment.get("speaker_confidence") or 0),
|
||||
"speaker_status": str(segment.get("speaker_status") or "confirmed"),
|
||||
}
|
||||
if usable and usable[-1]["speaker"]["speaker_id"] == speaker_id:
|
||||
usable[-1]["end_time_ms"] = end
|
||||
else:
|
||||
usable.append(
|
||||
{"start_time_ms": start, "end_time_ms": end, "speaker": speaker}
|
||||
)
|
||||
|
||||
if not usable:
|
||||
return []
|
||||
if len(usable) == 1:
|
||||
usable[0]["text"] = text
|
||||
return usable
|
||||
|
||||
# 双重保护:服务端有最短段长滤波,桥接层也拒绝意外的短片段响应。
|
||||
min_split_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_MIN_SEGMENT_MS", "3000")))
|
||||
while len(usable) > 1:
|
||||
short_index = next(
|
||||
(
|
||||
index
|
||||
for index, item in enumerate(usable)
|
||||
if item["end_time_ms"] - item["start_time_ms"] < min_split_ms
|
||||
),
|
||||
None,
|
||||
)
|
||||
if short_index is None:
|
||||
break
|
||||
if short_index == 0:
|
||||
usable[1]["start_time_ms"] = usable[0]["start_time_ms"]
|
||||
del usable[0]
|
||||
elif short_index == len(usable) - 1:
|
||||
usable[-2]["end_time_ms"] = usable[-1]["end_time_ms"]
|
||||
del usable[-1]
|
||||
else:
|
||||
previous = usable[short_index - 1]
|
||||
following = usable[short_index + 1]
|
||||
if previous["end_time_ms"] - previous["start_time_ms"] >= (
|
||||
following["end_time_ms"] - following["start_time_ms"]
|
||||
):
|
||||
previous["end_time_ms"] = usable[short_index]["end_time_ms"]
|
||||
del usable[short_index]
|
||||
else:
|
||||
following["start_time_ms"] = usable[short_index]["start_time_ms"]
|
||||
del usable[short_index]
|
||||
|
||||
if len(usable) == 1 or len(text) < len(usable):
|
||||
usable = [max(usable, key=lambda item: item["end_time_ms"] - item["start_time_ms"])]
|
||||
usable[0]["text"] = text
|
||||
return usable
|
||||
|
||||
total_duration = sum(item["end_time_ms"] - item["start_time_ms"] for item in usable)
|
||||
char_start = 0
|
||||
elapsed = 0.0
|
||||
for index, item in enumerate(usable):
|
||||
elapsed += item["end_time_ms"] - item["start_time_ms"]
|
||||
if index == len(usable) - 1:
|
||||
char_end = len(text)
|
||||
else:
|
||||
char_end = round(len(text) * elapsed / total_duration)
|
||||
char_end = max(char_start + 1, min(char_end, len(text) - (len(usable) - index - 1)))
|
||||
item["text"] = text[char_start:char_end].strip()
|
||||
char_start = char_end
|
||||
return [item for item in usable if item.get("text")]
|
||||
|
||||
|
||||
class BrowserSession:
|
||||
"""Own one browser/native WS pair and translate their message contracts."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
browser_ws: web.WebSocketResponse,
|
||||
auxiliary: AuxiliaryModelService,
|
||||
start: dict[str, Any],
|
||||
voice_id: str,
|
||||
) -> None:
|
||||
self.browser_ws = browser_ws
|
||||
self.auxiliary = auxiliary
|
||||
self.start = start
|
||||
self.voice_id = voice_id
|
||||
self.session_id = uuid4().hex
|
||||
self.stop_event = asyncio.Event()
|
||||
self.send_lock = asyncio.Lock()
|
||||
self.speaker_jobs: asyncio.Queue[SpeakerJob | None] = asyncio.Queue()
|
||||
self.wav_decoder = (
|
||||
IncrementalWavDecoder()
|
||||
if str(start.get("source") or "mic") == "file"
|
||||
and Path(str(start.get("file_name") or "")).suffix.lower() == ".wav"
|
||||
else None
|
||||
)
|
||||
self.speed_factor = max(0.5, min(3.0, float(start.get("speed_factor") or 1.0)))
|
||||
try:
|
||||
sentence_strategy = int(start.get("sentence_strategy", 0))
|
||||
except (TypeError, ValueError):
|
||||
sentence_strategy = 0
|
||||
self.sentence_strategy = sentence_strategy if sentence_strategy in (0, 1) else 0
|
||||
self.pending_pcm = bytearray()
|
||||
self.turn_audio = bytearray()
|
||||
self.total_audio_ms = 0.0
|
||||
self.turn_start_ms = 0.0
|
||||
self.turn_text = ""
|
||||
self.sentence_id = 0
|
||||
self.native_ack: dict[str, Any] | None = None
|
||||
self.native_error: str | None = None
|
||||
|
||||
async def emit(self, payload: dict[str, Any]) -> None:
|
||||
"""Serialize browser writes because ASR and CAM++ finish independently."""
|
||||
async with self.send_lock:
|
||||
if not self.browser_ws.closed:
|
||||
await self.browser_ws.send_json(payload)
|
||||
|
||||
async def emit_sentence(
|
||||
self,
|
||||
text: str,
|
||||
final: bool,
|
||||
speaker: dict[str, Any] | None = None,
|
||||
sentence_id: int | None = None,
|
||||
start_time_ms: float | None = None,
|
||||
end_time_ms: float | None = None,
|
||||
) -> None:
|
||||
speaker = speaker or {}
|
||||
sentence = {
|
||||
"sentence_id": self.sentence_id if sentence_id is None else sentence_id,
|
||||
"sentence": text,
|
||||
"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": int(speaker.get("speaker_id", -1)),
|
||||
"speaker_name": str(speaker.get("speaker_name") or ""),
|
||||
"speaker_confidence": float(speaker.get("speaker_confidence") or 0),
|
||||
"speaker_status": str(speaker.get("speaker_status") or "pending"),
|
||||
}
|
||||
await self.emit({"type": "sentences", "sentences": [sentence]})
|
||||
|
||||
def align_turn_audio_to_vad(self, message: dict[str, Any]) -> None:
|
||||
"""用原生 VAD 的绝对时间裁掉当前 turn 的前后静音样本。"""
|
||||
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
|
||||
try:
|
||||
speech_start_ms = message.get("speech_start_ms")
|
||||
if speech_start_ms is not None:
|
||||
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_start_ms)))
|
||||
trim_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
|
||||
trim_bytes -= trim_bytes % 2
|
||||
del self.turn_audio[:trim_bytes]
|
||||
self.turn_start_ms += trim_bytes / PCM_BYTES_PER_MS
|
||||
buffer_end_ms = self.turn_start_ms + len(self.turn_audio) / PCM_BYTES_PER_MS
|
||||
|
||||
speech_end_ms = message.get("speech_end_ms")
|
||||
if speech_end_ms is not None:
|
||||
target = max(self.turn_start_ms, min(buffer_end_ms, float(speech_end_ms)))
|
||||
keep_bytes = int(round((target - self.turn_start_ms) * PCM_BYTES_PER_MS))
|
||||
keep_bytes -= keep_bytes % 2
|
||||
del self.turn_audio[keep_bytes:]
|
||||
except (TypeError, ValueError):
|
||||
LOGGER.warning("Ignoring invalid native FunASR VAD boundary: %s", message)
|
||||
|
||||
async def send_pcm_frame(self, native_ws: Any, frame: bytes, pace_file: bool) -> None:
|
||||
if not frame:
|
||||
return
|
||||
if len(frame) % 2:
|
||||
raise ValueError("PCM16 audio ended on an incomplete sample")
|
||||
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.
|
||||
trim = len(self.turn_audio) - MAX_SPEAKER_AUDIO_BYTES
|
||||
del self.turn_audio[:trim]
|
||||
self.turn_start_ms += trim / PCM_BYTES_PER_MS
|
||||
await native_ws.send(frame)
|
||||
if pace_file:
|
||||
await asyncio.sleep(len(frame) / (SAMPLE_RATE * 2) / self.speed_factor)
|
||||
|
||||
async def accept_audio(self, native_ws: Any, data: bytes) -> None:
|
||||
pcm = self.wav_decoder.feed(data) if self.wav_decoder else data
|
||||
if not pcm:
|
||||
return
|
||||
self.pending_pcm.extend(pcm)
|
||||
is_file = str(self.start.get("source") or "mic") == "file"
|
||||
while len(self.pending_pcm) >= FRAME_BYTES:
|
||||
frame = bytes(self.pending_pcm[:FRAME_BYTES])
|
||||
del self.pending_pcm[:FRAME_BYTES]
|
||||
await self.send_pcm_frame(native_ws, frame, is_file)
|
||||
|
||||
async def finish_audio(self, native_ws: Any) -> None:
|
||||
if self.wav_decoder:
|
||||
self.wav_decoder.finish()
|
||||
if self.pending_pcm:
|
||||
frame = bytes(self.pending_pcm)
|
||||
self.pending_pcm.clear()
|
||||
await self.send_pcm_frame(
|
||||
native_ws, frame, str(self.start.get("source") or "mic") == "file"
|
||||
)
|
||||
|
||||
async def read_native(self, native_ws: Any) -> None:
|
||||
"""Consume native FunASR events and keep its per-utterance partial cache."""
|
||||
try:
|
||||
while True:
|
||||
raw = await native_ws.recv()
|
||||
message = json.loads(raw)
|
||||
if message.get("is_end"):
|
||||
self.native_ack = message
|
||||
return
|
||||
if message.get("event") == "vad":
|
||||
self.align_turn_audio_to_vad(message)
|
||||
continue
|
||||
text = str(message.get("text") or "")
|
||||
if text:
|
||||
# FunASR online sends the newly decoded text for each chunk.
|
||||
self.turn_text += text
|
||||
if message.get("is_final"):
|
||||
final_text = self.turn_text.strip()
|
||||
audio = bytes(self.turn_audio)
|
||||
start_ms = self.turn_start_ms
|
||||
end_ms = start_ms + len(audio) / PCM_BYTES_PER_MS
|
||||
turn_sentence_id = self.sentence_id
|
||||
if final_text:
|
||||
# 先结束前端 interim 气泡;标点和 CAM++ 在独立 worker 中完成,
|
||||
# 不阻塞 native WS 继续读取后续音频帧。
|
||||
await self.emit_sentence(
|
||||
final_text,
|
||||
final=True,
|
||||
sentence_id=turn_sentence_id,
|
||||
start_time_ms=start_ms,
|
||||
end_time_ms=end_ms,
|
||||
)
|
||||
await self.speaker_jobs.put(
|
||||
SpeakerJob(
|
||||
sentence_id=turn_sentence_id,
|
||||
text=final_text,
|
||||
audio=audio,
|
||||
start_time_ms=start_ms,
|
||||
end_time_ms=end_ms,
|
||||
)
|
||||
)
|
||||
# 为同一个 VAD turn 内可能拆出的多个气泡预留独立 ID。
|
||||
self.sentence_id += TURN_SENTENCE_ID_STRIDE
|
||||
self.turn_text = ""
|
||||
self.turn_audio.clear()
|
||||
self.turn_start_ms = self.total_audio_ms
|
||||
elif self.turn_text:
|
||||
await self.emit_sentence(self.turn_text, final=False)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
self.native_error = str(exc)
|
||||
LOGGER.exception("native FunASR WebSocket closed unexpectedly")
|
||||
await self.emit({"type": "error", "message": f"FunASR realtime WS: {exc}"})
|
||||
|
||||
async def resolve_speakers(self) -> None:
|
||||
"""按序定稿标点并用 CAM++ 滑窗恢复 turn 内说话人切换。"""
|
||||
while True:
|
||||
job = await self.speaker_jobs.get()
|
||||
try:
|
||||
if job is None:
|
||||
return
|
||||
final_text = job.text
|
||||
try:
|
||||
punctuation = await self.auxiliary.punctuate(final_text)
|
||||
if not punctuation.get("available", False):
|
||||
LOGGER.warning(
|
||||
"FunASR punctuation is unavailable: %s",
|
||||
punctuation.get("error", "model is not configured"),
|
||||
)
|
||||
else:
|
||||
punctuated = str(punctuation.get("text") or "").strip()
|
||||
if punctuated:
|
||||
final_text = punctuated
|
||||
except Exception:
|
||||
LOGGER.exception("FunASR punctuation request failed; keeping raw text")
|
||||
|
||||
subsegments: list[dict[str, Any]] = []
|
||||
if len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
|
||||
# required CAM++ 按 FunASR 1.5s/0.75s 滑窗识别同一 VAD turn 内的
|
||||
# 多人切换;若窗长不足或接口暂不可用,再退回整段声纹验证。
|
||||
track = getattr(self.auxiliary, "track_speakers", None)
|
||||
if callable(track):
|
||||
try:
|
||||
tracked = await track(
|
||||
job.audio,
|
||||
self.session_id,
|
||||
job.start_time_ms,
|
||||
job.end_time_ms,
|
||||
)
|
||||
subsegments = split_text_by_speaker_segments(
|
||||
final_text,
|
||||
tracked,
|
||||
job.start_time_ms,
|
||||
job.end_time_ms,
|
||||
)
|
||||
except Exception:
|
||||
LOGGER.exception(
|
||||
"CAM++ sliding-window diarization failed; falling back to whole-turn speaker"
|
||||
)
|
||||
|
||||
if not subsegments and len(job.audio) >= MIN_SPEAKER_AUDIO_BYTES:
|
||||
try:
|
||||
resolved = await self.auxiliary.resolve_speaker(
|
||||
job.audio,
|
||||
self.session_id,
|
||||
job.start_time_ms,
|
||||
job.end_time_ms,
|
||||
)
|
||||
if resolved and int(resolved.get("speaker_id", -1)) >= 0:
|
||||
subsegments = [
|
||||
{
|
||||
"text": final_text,
|
||||
"speaker": resolved,
|
||||
"start_time_ms": job.start_time_ms,
|
||||
"end_time_ms": job.end_time_ms,
|
||||
}
|
||||
]
|
||||
except Exception:
|
||||
LOGGER.exception("CAM++ speaker resolution failed: voice_id=%s", self.voice_id)
|
||||
|
||||
if not subsegments:
|
||||
subsegments = [
|
||||
{
|
||||
"text": final_text,
|
||||
"speaker": {"speaker_id": -1},
|
||||
"start_time_ms": job.start_time_ms,
|
||||
"end_time_ms": job.end_time_ms,
|
||||
}
|
||||
]
|
||||
|
||||
for index, segment in enumerate(subsegments):
|
||||
await self.emit_sentence(
|
||||
str(segment["text"]),
|
||||
final=True,
|
||||
speaker=segment.get("speaker"),
|
||||
sentence_id=job.sentence_id + index,
|
||||
start_time_ms=float(segment.get("start_time_ms", job.start_time_ms)),
|
||||
end_time_ms=float(segment.get("end_time_ms", job.end_time_ms)),
|
||||
)
|
||||
finally:
|
||||
self.speaker_jobs.task_done()
|
||||
|
||||
|
||||
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",
|
||||
"model": os.getenv("FUNASR_ASR_MODEL", ""),
|
||||
"native_ws_url": NATIVE_WS_URL,
|
||||
"speaker_service_url": os.getenv(
|
||||
"AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def stop_handler(request: web.Request) -> web.Response:
|
||||
voice_id = request.query.get("voice_id", "").strip()
|
||||
if not voice_id:
|
||||
return web.json_response({"ok": False, "error": "missing voice_id"}, status=400)
|
||||
stop_event = request.app[SESSION_REGISTRY_KEY].get(voice_id)
|
||||
if stop_event is None:
|
||||
return web.json_response({"ok": False, "error": "session not found"})
|
||||
stop_event.set()
|
||||
return web.json_response({"ok": True})
|
||||
|
||||
|
||||
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
|
||||
"""Bridge Tencent's browser messages to FunASR's native realtime protocol."""
|
||||
browser_ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024, heartbeat=30)
|
||||
await browser_ws.prepare(request)
|
||||
session: BrowserSession | None = None
|
||||
native_reader: asyncio.Task[None] | None = None
|
||||
speaker_worker: asyncio.Task[None] | None = None
|
||||
voice_id = ""
|
||||
registered = False
|
||||
|
||||
try:
|
||||
first = await browser_ws.receive()
|
||||
if first.type != WSMsgType.TEXT:
|
||||
await browser_ws.send_json({"type": "error", "message": "first message must be JSON start"})
|
||||
return browser_ws
|
||||
start = json.loads(first.data)
|
||||
if not isinstance(start, dict) or start.get("type") != "start":
|
||||
await browser_ws.send_json({"type": "error", "message": "first message must have type=start"})
|
||||
return browser_ws
|
||||
source = str(start.get("source") or "mic")
|
||||
suffix = Path(str(start.get("file_name") or "")).suffix.lower()
|
||||
if source == "file" and suffix not in {".pcm", ".wav"}:
|
||||
await browser_ws.send_json(
|
||||
{
|
||||
"type": "error",
|
||||
"message": "FunASR realtime mode accepts PCM or 16 kHz mono PCM WAV files.",
|
||||
}
|
||||
)
|
||||
return browser_ws
|
||||
if source == "file" and suffix == ".pcm":
|
||||
start = dict(start)
|
||||
start["file_name"] = str(start.get("file_name") or "audio.pcm")
|
||||
|
||||
voice_id = uuid4().hex
|
||||
auxiliary: AuxiliaryModelService = request.app[AUXILIARY_KEY]
|
||||
session = BrowserSession(browser_ws, auxiliary, start, voice_id)
|
||||
request.app[SESSION_REGISTRY_KEY][voice_id] = session.stop_event
|
||||
registered = True
|
||||
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.
|
||||
async with websocket_connect(
|
||||
NATIVE_WS_URL,
|
||||
subprotocols=["binary"],
|
||||
ping_interval=None,
|
||||
close_timeout=3,
|
||||
max_size=None,
|
||||
) as native_ws:
|
||||
await native_ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"mode": "online",
|
||||
"chunk_size": list(CHUNK_SIZE),
|
||||
"chunk_interval": CHUNK_INTERVAL,
|
||||
"encoder_chunk_look_back": int(
|
||||
os.getenv("FUNASR_ENCODER_LOOK_BACK", "4")
|
||||
),
|
||||
"decoder_chunk_look_back": int(
|
||||
os.getenv("FUNASR_DECODER_LOOK_BACK", "1")
|
||||
),
|
||||
"sentence_strategy": session.sentence_strategy,
|
||||
"audio_fs": SAMPLE_RATE,
|
||||
"wav_name": voice_id,
|
||||
"is_speaking": True,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
)
|
||||
native_reader = asyncio.create_task(session.read_native(native_ws))
|
||||
speaker_worker = asyncio.create_task(session.resolve_speakers())
|
||||
|
||||
while not browser_ws.closed:
|
||||
receive_task = asyncio.create_task(browser_ws.receive())
|
||||
stop_task = asyncio.create_task(session.stop_event.wait())
|
||||
done, _ = await asyncio.wait(
|
||||
[receive_task, stop_task, native_reader],
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if stop_task in done:
|
||||
receive_task.cancel()
|
||||
await asyncio.gather(receive_task, return_exceptions=True)
|
||||
break
|
||||
stop_task.cancel()
|
||||
await asyncio.gather(stop_task, return_exceptions=True)
|
||||
if native_reader in done:
|
||||
receive_task.cancel()
|
||||
await asyncio.gather(receive_task, return_exceptions=True)
|
||||
if session.native_ack is None and session.native_error is None:
|
||||
session.native_error = "FunASR native WebSocket ended before EOF acknowledgement"
|
||||
break
|
||||
|
||||
message = await receive_task
|
||||
if message.type == WSMsgType.BINARY:
|
||||
await session.accept_audio(native_ws, bytes(message.data))
|
||||
continue
|
||||
if message.type == WSMsgType.TEXT:
|
||||
try:
|
||||
control = json.loads(message.data)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(control, dict) and control.get("type") in {"eof", "stop"}:
|
||||
break
|
||||
if isinstance(control, dict) and control.get("type") == "abort":
|
||||
session.native_error = "session aborted by browser"
|
||||
break
|
||||
if message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}:
|
||||
break
|
||||
|
||||
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.
|
||||
await native_ws.send(
|
||||
json.dumps({"is_speaking": False, "is_end": True}, ensure_ascii=False)
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(native_reader, timeout=FINALIZE_TIMEOUT_SECONDS)
|
||||
except asyncio.TimeoutError:
|
||||
session.native_error = (
|
||||
f"FunASR did not acknowledge end-of-input within "
|
||||
f"{FINALIZE_TIMEOUT_SECONDS}s"
|
||||
)
|
||||
|
||||
if speaker_worker is not None:
|
||||
await session.speaker_jobs.join()
|
||||
await session.speaker_jobs.put(None)
|
||||
await speaker_worker
|
||||
speaker_worker = None
|
||||
if session.native_error:
|
||||
await session.emit({"type": "error", "message": session.native_error})
|
||||
elif session.native_ack and not session.native_ack.get("is_final", False):
|
||||
await session.emit(
|
||||
{
|
||||
"type": "error",
|
||||
"message": str(session.native_ack.get("error") or "FunASR did not finalize the stream"),
|
||||
}
|
||||
)
|
||||
elif session.native_ack:
|
||||
await session.emit({"type": "end"})
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
LOGGER.exception("Tencent-compatible WebSocket session failed")
|
||||
if not browser_ws.closed:
|
||||
await browser_ws.send_json({"type": "error", "message": str(exc)})
|
||||
finally:
|
||||
if native_reader is not None and not native_reader.done():
|
||||
native_reader.cancel()
|
||||
await asyncio.gather(native_reader, return_exceptions=True)
|
||||
if speaker_worker is not None and not speaker_worker.done():
|
||||
speaker_worker.cancel()
|
||||
await asyncio.gather(speaker_worker, return_exceptions=True)
|
||||
if registered:
|
||||
request.app[SESSION_REGISTRY_KEY].pop(voice_id, None)
|
||||
if session is not None:
|
||||
await session.auxiliary.reset_speaker_session(session.session_id)
|
||||
if not browser_ws.closed:
|
||||
await browser_ws.close()
|
||||
return browser_ws
|
||||
|
||||
|
||||
async def create_app() -> web.Application:
|
||||
"""Create a light protocol bridge; model inference belongs to native FunASR."""
|
||||
app = web.Application()
|
||||
app[AUXILIARY_KEY] = AuxiliaryModelService(
|
||||
AuxiliaryServiceConfig(
|
||||
base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010")
|
||||
)
|
||||
)
|
||||
app[SESSION_REGISTRY_KEY] = {}
|
||||
|
||||
async def lifecycle(application: web.Application):
|
||||
await application[AUXILIARY_KEY].start()
|
||||
try:
|
||||
health = await application[AUXILIARY_KEY].health()
|
||||
if not health.get("speaker_embedding_ready"):
|
||||
raise RuntimeError("CAM++ speaker service is not ready")
|
||||
except Exception:
|
||||
await application[AUXILIARY_KEY].close()
|
||||
raise
|
||||
yield
|
||||
await application[AUXILIARY_KEY].close()
|
||||
|
||||
app.cleanup_ctx.append(lifecycle)
|
||||
app.router.add_get("/api/config", config_handler)
|
||||
app.router.add_get("/api/stop", stop_handler)
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
return app
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--no-browser", action="store_true")
|
||||
args = parser.parse_args()
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
print(
|
||||
f"FunASR browser bridge: http://{os.getenv('WEB_DISPLAY_HOST', '127.0.0.1')}:{WEB_PORT}/api/config",
|
||||
flush=True,
|
||||
)
|
||||
web.run_app(create_app(), host=WEB_HOST, port=WEB_PORT)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,75 +0,0 @@
|
|||
"""Unit tests for the migrated FunASR streaming lifecycle."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
try:
|
||||
from backend.realtime_websocket.funasr_engine import (
|
||||
FunASRRealtimeSession,
|
||||
FunASRServiceConfig,
|
||||
_result_text,
|
||||
_vad_events,
|
||||
)
|
||||
except ModuleNotFoundError:
|
||||
from funasr_engine import (
|
||||
FunASRRealtimeSession,
|
||||
FunASRServiceConfig,
|
||||
_result_text,
|
||||
_vad_events,
|
||||
)
|
||||
|
||||
|
||||
class FakeService:
|
||||
def __init__(self) -> None:
|
||||
self.config = FunASRServiceConfig(
|
||||
device="cpu",
|
||||
vad_device="cpu",
|
||||
vad_chunk_ms=1000,
|
||||
chunk_size=(0, 2, 1),
|
||||
)
|
||||
self.vad_calls = 0
|
||||
self.asr_calls: list[dict[str, object]] = []
|
||||
|
||||
async def generate_vad(self, audio, status, chunk_ms):
|
||||
self.vad_calls += 1
|
||||
return [(0, -1)] if self.vad_calls == 1 else [(-1, 2000)]
|
||||
|
||||
async def generate_asr(self, audio, status):
|
||||
self.asr_calls.append(status)
|
||||
return "第一块" if len(self.asr_calls) == 1 else "第二块"
|
||||
|
||||
|
||||
class FunASREngineTests(unittest.TestCase):
|
||||
def test_result_normalization(self) -> None:
|
||||
self.assertEqual(_result_text([{"text": "你好"}]), "你好")
|
||||
self.assertEqual(_result_text({"value": "片段"}), "片段")
|
||||
self.assertEqual(
|
||||
_vad_events([{"value": [[0, -1], [-1, 800]]}]),
|
||||
[(0.0, -1.0), (-1.0, 800.0)],
|
||||
)
|
||||
|
||||
def test_stream_has_independent_cache_and_final_flush(self) -> None:
|
||||
async def exercise() -> None:
|
||||
service = FakeService()
|
||||
with patch.object(FunASRRealtimeSession, "_to_float32", staticmethod(lambda value: value)):
|
||||
first = FunASRRealtimeSession(service)
|
||||
second = FunASRRealtimeSession(service)
|
||||
first_events = await first.feed(b"\\x01\\x00" * 16000)
|
||||
second_events = await second.feed(b"\\x01\\x00" * 16000)
|
||||
first_events += await first.feed(b"\\x01\\x00" * 16000)
|
||||
first_events += await first.finish()
|
||||
self.assertTrue(any(not event.is_final for event in first_events))
|
||||
self.assertTrue(any(event.is_final for event in first_events))
|
||||
self.assertEqual(first.segment_id, 1)
|
||||
self.assertEqual(second.segment_id, 0)
|
||||
self.assertTrue(any(status["is_final"] is True for status in service.asr_calls))
|
||||
self.assertGreaterEqual(len(service.asr_calls), 2)
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -1,127 +0,0 @@
|
|||
"""Contract test for the unchanged Tencent UI to FunASR native WS bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import AioHTTPTestCase
|
||||
|
||||
from backend.realtime_websocket.funasr_server import (
|
||||
AUXILIARY_KEY,
|
||||
SESSION_REGISTRY_KEY,
|
||||
websocket_handler,
|
||||
)
|
||||
|
||||
|
||||
class FakeNativeWebSocket:
|
||||
"""Stand in for FunASR's native WSS process without loading model weights."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.incoming: asyncio.Queue[str] = asyncio.Queue()
|
||||
self.audio_bytes = 0
|
||||
self.config: dict[str, object] = {}
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *_args):
|
||||
return None
|
||||
|
||||
async def send(self, payload: str | bytes) -> None:
|
||||
if isinstance(payload, bytes):
|
||||
self.audio_bytes += len(payload)
|
||||
return
|
||||
control = json.loads(payload)
|
||||
if "mode" in control:
|
||||
self.config = control
|
||||
return
|
||||
if control.get("is_end"):
|
||||
await self.incoming.put(
|
||||
json.dumps({"mode": "online", "text": "hello", "is_final": False})
|
||||
)
|
||||
await self.incoming.put(
|
||||
json.dumps({"mode": "online", "text": " world", "is_final": True})
|
||||
)
|
||||
await self.incoming.put(
|
||||
json.dumps({"is_end": True, "is_final": True})
|
||||
)
|
||||
|
||||
async def recv(self) -> str:
|
||||
return await self.incoming.get()
|
||||
|
||||
|
||||
class FakeAuxiliaryService:
|
||||
async def resolve_speaker(self, _audio, _session_id, _start, _end):
|
||||
return {
|
||||
"speaker_id": 0,
|
||||
"speaker_name": "speaker 1",
|
||||
"speaker_confidence": 0.9,
|
||||
"speaker_status": "confirmed",
|
||||
}
|
||||
|
||||
async def reset_speaker_session(self, _session_id):
|
||||
return None
|
||||
|
||||
|
||||
class FunASRBridgeTests(AioHTTPTestCase):
|
||||
def get_app(self):
|
||||
app = web.Application()
|
||||
app[AUXILIARY_KEY] = FakeAuxiliaryService()
|
||||
app[SESSION_REGISTRY_KEY] = {}
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
return app
|
||||
|
||||
async def test_tencent_ui_messages_use_native_funasr_and_keep_speaker_label(self):
|
||||
native = FakeNativeWebSocket()
|
||||
with patch(
|
||||
"backend.realtime_websocket.funasr_server.websocket_connect",
|
||||
return_value=native,
|
||||
):
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json(
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
# The server keeps speaker labeling enabled even if this flag is false.
|
||||
"speaker_diarization": 0,
|
||||
}
|
||||
)
|
||||
first = await ws.receive_json()
|
||||
second = await ws.receive_json()
|
||||
self.assertEqual(first["type"], "voice_id")
|
||||
self.assertEqual(second["type"], "start")
|
||||
|
||||
pcm = b"\x01\x00" * 16000
|
||||
await ws.send_bytes(pcm)
|
||||
await ws.send_json({"type": "eof"})
|
||||
|
||||
messages = []
|
||||
async with asyncio.timeout(5):
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
messages.append(message)
|
||||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
self.assertEqual(native.config["mode"], "online")
|
||||
self.assertEqual(native.config["audio_fs"], 16000)
|
||||
self.assertEqual(native.audio_bytes, len(pcm))
|
||||
sentence_events = [
|
||||
sentence
|
||||
for message in messages
|
||||
if message["type"] == "sentences"
|
||||
for sentence in message["sentences"]
|
||||
]
|
||||
self.assertTrue(any(sentence["sentence_type"] == 0 for sentence in sentence_events))
|
||||
final_events = [sentence for sentence in sentence_events if sentence["sentence_type"] == 1]
|
||||
self.assertEqual(final_events[-1]["sentence"], "hello world")
|
||||
self.assertEqual(final_events[-1]["speaker_id"], 0)
|
||||
await ws.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -1,261 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start local CAM++, FunASR native realtime WSS, and the browser protocol bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
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
|
||||
|
||||
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.
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from backend.model_manifest import (
|
||||
auxiliary_models,
|
||||
load_manifest,
|
||||
model_directory,
|
||||
resolve_auxiliary_model_id,
|
||||
resolve_model_id,
|
||||
)
|
||||
|
||||
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."""
|
||||
name = requested.strip()
|
||||
configured_path = Path(name)
|
||||
direct_candidates = (
|
||||
[configured_path]
|
||||
if configured_path.is_absolute()
|
||||
else [PROJECT_ROOT / configured_path, models_dir / configured_path]
|
||||
)
|
||||
for candidate in direct_candidates:
|
||||
if candidate.is_dir() and (candidate / "configuration.json").is_file():
|
||||
return candidate.resolve()
|
||||
|
||||
manifest = load_manifest()
|
||||
if kind == "asr":
|
||||
model_id = resolve_model_id(name, manifest)
|
||||
else:
|
||||
model_id = resolve_auxiliary_model_id(name, manifest, kind=kind)
|
||||
candidate = model_directory(model_id, manifest, models_dir)
|
||||
if candidate.is_dir() and (candidate / "configuration.json").is_file():
|
||||
return candidate.resolve()
|
||||
raise FileNotFoundError(
|
||||
f"Local {kind.upper()} model '{name}' is missing; expected: {candidate}"
|
||||
)
|
||||
|
||||
|
||||
def local_cam_model(models_dir: Path) -> Path:
|
||||
"""Require a complete CAM++ speaker verification asset from the manifest."""
|
||||
manifest = load_manifest()
|
||||
override = os.getenv("CAM_MODEL_PATH", "").strip()
|
||||
if override:
|
||||
path = Path(override)
|
||||
if not path.is_absolute():
|
||||
path = models_dir / path
|
||||
path = path.resolve()
|
||||
configs = [
|
||||
config for config in auxiliary_models(manifest).values()
|
||||
if config.get("kind") == "speaker_verification"
|
||||
]
|
||||
if path.is_dir() and configs and all(
|
||||
(path / relative).is_file() for relative in configs[0].get("required_files", [])
|
||||
):
|
||||
return path
|
||||
raise FileNotFoundError(f"CAM++ model is missing or incomplete: {path}")
|
||||
checked = []
|
||||
for model_id, config in auxiliary_models(manifest).items():
|
||||
if config.get("kind") != "speaker_verification":
|
||||
continue
|
||||
path = model_directory(model_id, manifest, models_dir)
|
||||
required = [path / relative for relative in config.get("required_files", [])]
|
||||
if path.is_dir() and all(item.is_file() for item in required):
|
||||
return path.resolve()
|
||||
checked.append(str(path))
|
||||
raise FileNotFoundError("CAM++ speaker model is missing; checked: " + ", ".join(checked))
|
||||
|
||||
|
||||
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."""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"Service exited before it became ready ({url}, exit={code})")
|
||||
try:
|
||||
with urlopen(url, timeout=2) as response:
|
||||
data = json.load(response)
|
||||
if data.get(key):
|
||||
return
|
||||
except (OSError, ValueError, URLError):
|
||||
pass
|
||||
time.sleep(0.5)
|
||||
raise TimeoutError(f"Service did not become ready within {seconds}s: {url}")
|
||||
|
||||
|
||||
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."""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"FunASR native WS exited before readiness (exit={code})")
|
||||
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:
|
||||
raise RuntimeError(f"FunASR native WS exited before readiness (exit={code})")
|
||||
return
|
||||
except OSError:
|
||||
time.sleep(0.5)
|
||||
raise TimeoutError(f"FunASR native WS did not open {host}:{port} within {seconds}s")
|
||||
|
||||
|
||||
def stop_child(process: subprocess.Popen | None) -> None:
|
||||
"""Stop a supervised model or WebSocket process on launcher shutdown."""
|
||||
if process is None or process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Require ASR, VAD, and CAM++ before exposing the public WS bridge."""
|
||||
models_dir = Path(os.getenv("MODEL_DIR", "models"))
|
||||
if not models_dir.is_absolute():
|
||||
models_dir = PROJECT_ROOT / models_dir
|
||||
models_dir = models_dir.resolve()
|
||||
asr = local_model(
|
||||
os.getenv("FUNASR_ASR_MODEL", "paraformer-zh-streaming"), models_dir, "asr"
|
||||
)
|
||||
vad = local_model(os.getenv("FUNASR_VAD_MODEL", "fsmn-vad"), models_dir, "vad")
|
||||
cam = local_cam_model(models_dir)
|
||||
|
||||
native_host = os.getenv("FUNASR_NATIVE_WS_HOST", "127.0.0.1")
|
||||
native_port = int(os.getenv("FUNASR_NATIVE_WS_PORT", "10095"))
|
||||
native_url = f"ws://{native_host}:{native_port}"
|
||||
aux_url = os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010").rstrip("/")
|
||||
device = os.getenv("FUNASR_DEVICE", "cuda:0")
|
||||
vad_device = os.getenv("FUNASR_VAD_DEVICE", "cpu")
|
||||
ngpu = "1" if device.startswith("cuda") else "0"
|
||||
ncpu = os.getenv("FUNASR_NCPU", str(os.cpu_count() or 4))
|
||||
web_port = int(os.getenv("WEB_PORT", "8082"))
|
||||
|
||||
print(f"Local models: ASR={asr}; VAD={vad}; CAM++={cam}", flush=True)
|
||||
print(
|
||||
f"Runtime: ASR device={device}; VAD device={vad_device}; native WS={native_url}",
|
||||
flush=True,
|
||||
)
|
||||
env = os.environ.copy()
|
||||
preload_kinds = {
|
||||
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.
|
||||
preload_kinds.discard("vad")
|
||||
preload_kinds.add("speaker_verification")
|
||||
env.update(
|
||||
{
|
||||
"AUXILIARY_PRELOAD_KINDS": ",".join(sorted(preload_kinds)),
|
||||
"MODEL_DIR": str(models_dir),
|
||||
"FUNASR_ASR_MODEL": str(asr),
|
||||
"FUNASR_VAD_MODEL": str(vad),
|
||||
"FUNASR_DEVICE": device,
|
||||
"FUNASR_VAD_DEVICE": vad_device,
|
||||
"FUNASR_NATIVE_WS_URL": native_url,
|
||||
"CAM_MODEL_PATH": str(cam),
|
||||
"AUXILIARY_SERVICE_URL": aux_url,
|
||||
}
|
||||
)
|
||||
|
||||
auxiliary = None
|
||||
native = None
|
||||
websocket = None
|
||||
try:
|
||||
auxiliary = subprocess.Popen(
|
||||
[sys.executable, "-m", "backend.auxiliary_server"],
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
wait_for_health(f"{aux_url}/health", auxiliary, "speaker_embedding_ready")
|
||||
|
||||
native_args = [
|
||||
sys.executable,
|
||||
str(PROJECT_ROOT / "backend" / "realtime_websocket" / "funasr_native_wss.py"),
|
||||
"--host",
|
||||
native_host,
|
||||
"--port",
|
||||
str(native_port),
|
||||
"--asr_model",
|
||||
"",
|
||||
"--asr_model_online",
|
||||
str(asr),
|
||||
"--vad_model",
|
||||
str(vad),
|
||||
"--punc_model",
|
||||
"",
|
||||
"--device",
|
||||
device,
|
||||
"--vad_device",
|
||||
vad_device,
|
||||
"--ngpu",
|
||||
ngpu,
|
||||
"--ncpu",
|
||||
ncpu,
|
||||
"--certfile",
|
||||
"",
|
||||
"--keyfile",
|
||||
"",
|
||||
]
|
||||
native = subprocess.Popen(native_args, cwd=PROJECT_ROOT, env=env)
|
||||
probe_host = "127.0.0.1" if native_host in {"0.0.0.0", "::"} else native_host
|
||||
wait_for_tcp(probe_host, native_port, native)
|
||||
|
||||
websocket = subprocess.Popen(
|
||||
[sys.executable, "-m", "backend.run_funasr_demo", "--no-browser"],
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
wait_for_health(
|
||||
f"http://127.0.0.1:{web_port}/api/config", websocket, "engine"
|
||||
)
|
||||
print(f"FunASR browser backend ready: ws://127.0.0.1:{web_port}/ws", flush=True)
|
||||
while True:
|
||||
for label, process in (
|
||||
("CAM++", auxiliary),
|
||||
("FunASR native WS", native),
|
||||
("browser WS adapter", websocket),
|
||||
):
|
||||
code = process.poll()
|
||||
if code is not None:
|
||||
raise RuntimeError(f"{label} service exited (exit={code})")
|
||||
time.sleep(0.5)
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
finally:
|
||||
stop_child(websocket)
|
||||
stop_child(native)
|
||||
stop_child(auxiliary)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,16 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start the FunASR-backed browser demo."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from backend.realtime_websocket.funasr_server import main
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1,132 +0,0 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Serve the unchanged Tencent demo UI and proxy its API to the backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlsplit
|
||||
from typing import Any
|
||||
|
||||
from aiohttp import ClientSession, ClientTimeout, WSMsgType, web
|
||||
from dotenv import load_dotenv
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
STATIC_ROOT = Path(__file__).resolve().parent / "static"
|
||||
load_dotenv(PROJECT_ROOT / ".env")
|
||||
FRONTEND_HOST = os.getenv("FRONTEND_HOST", "127.0.0.1")
|
||||
FRONTEND_PORT = int(os.getenv("FRONTEND_PORT", "8080"))
|
||||
BACKEND_BASE_URL = os.getenv(
|
||||
"BACKEND_INTERNAL_URL",
|
||||
f"http://127.0.0.1:{os.getenv('WEB_PORT', '8082')}",
|
||||
).rstrip("/")
|
||||
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}"
|
||||
|
||||
|
||||
async def index_handler(_: web.Request) -> web.FileResponse:
|
||||
return web.FileResponse(
|
||||
STATIC_ROOT / "index.html", headers={"Cache-Control": "no-store"}
|
||||
)
|
||||
|
||||
|
||||
async def api_stop_proxy(request: web.Request) -> web.Response:
|
||||
"""Forward the Tencent page's existing stop request to the WS backend."""
|
||||
async with request.app[HTTP_SESSION].get(
|
||||
backend_url(request), timeout=ClientTimeout(total=5)
|
||||
) as response:
|
||||
body = await response.read()
|
||||
return web.Response(
|
||||
status=response.status,
|
||||
body=body,
|
||||
headers={"Content-Type": response.headers.get("Content-Type", "application/json")},
|
||||
)
|
||||
|
||||
|
||||
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:
|
||||
backend_ws = await request.app[HTTP_SESSION].ws_connect(
|
||||
backend_url(request),
|
||||
max_msg_size=64 * 1024 * 1024,
|
||||
heartbeat=30,
|
||||
autoping=True,
|
||||
)
|
||||
except Exception as exc:
|
||||
await browser_ws.send_json({"type": "error", "message": f"backend unavailable: {exc}"})
|
||||
await browser_ws.close(code=1011, message=b"backend unavailable")
|
||||
return browser_ws
|
||||
|
||||
async def relay(source: Any, destination: Any) -> None:
|
||||
async for message in source:
|
||||
if message.type == WSMsgType.TEXT:
|
||||
await destination.send_str(message.data)
|
||||
elif message.type == WSMsgType.BINARY:
|
||||
await destination.send_bytes(message.data)
|
||||
elif message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.ERROR}:
|
||||
if not destination.closed:
|
||||
code = message.data if isinstance(message.data, int) else 1000
|
||||
reason = message.extra or ""
|
||||
await destination.close(code=code, message=str(reason).encode("utf-8"))
|
||||
return
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(relay(browser_ws, backend_ws)),
|
||||
asyncio.create_task(relay(backend_ws, browser_ws)),
|
||||
]
|
||||
try:
|
||||
done, pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
await asyncio.gather(*done, *pending, return_exceptions=True)
|
||||
finally:
|
||||
if not backend_ws.closed:
|
||||
await backend_ws.close()
|
||||
if not browser_ws.closed:
|
||||
await browser_ws.close()
|
||||
return browser_ws
|
||||
|
||||
|
||||
async def create_app() -> web.Application:
|
||||
if not STATIC_ROOT.is_dir():
|
||||
raise FileNotFoundError(f"Tencent demo static directory is missing: {STATIC_ROOT}")
|
||||
parsed = urlsplit(BACKEND_BASE_URL)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc or parsed.path:
|
||||
raise ValueError("BACKEND_INTERNAL_URL must contain only an HTTP origin")
|
||||
|
||||
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)
|
||||
)
|
||||
yield
|
||||
await application[HTTP_SESSION].close()
|
||||
|
||||
app.cleanup_ctx.append(lifecycle)
|
||||
app.router.add_get("/", index_handler)
|
||||
app.router.add_get("/ws", websocket_proxy)
|
||||
app.router.add_get("/api/stop", api_stop_proxy)
|
||||
app.router.add_static("/", STATIC_ROOT, show_index=False)
|
||||
return app
|
||||
|
||||
|
||||
def main() -> None:
|
||||
print(
|
||||
f"Tencent demo frontend: http://{FRONTEND_HOST}:{FRONTEND_PORT}/ "
|
||||
f"(backend proxy: {BACKEND_BASE_URL})",
|
||||
flush=True,
|
||||
)
|
||||
web.run_app(create_app(), host=FRONTEND_HOST, port=FRONTEND_PORT)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -10,24 +10,6 @@
|
|||
"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",
|
||||
"description": "FunASR Paraformer streaming ASR used by the realtime backend",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"config.yaml"
|
||||
],
|
||||
"any_files": [
|
||||
"*.pb",
|
||||
"*.pt",
|
||||
"*.onnx",
|
||||
"*.bin",
|
||||
"*.model",
|
||||
"*.safetensors"
|
||||
],
|
||||
"min_total_size_bytes": 1000000
|
||||
}
|
||||
},
|
||||
"auxiliary_models": {
|
||||
|
|
@ -37,28 +19,7 @@
|
|||
"kind": "vad",
|
||||
"description": "FunASR FSMN VAD",
|
||||
"revision": "v2.0.2",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"config.yaml",
|
||||
"model.pb"
|
||||
],
|
||||
"min_total_size_bytes": 1000000,
|
||||
"funasr_alias": "fsmn-vad"
|
||||
},
|
||||
"iic/punc_ct-transformer_zh-cn-common-vocab272727-pytorch": {
|
||||
"alias": "ct-punc",
|
||||
"directory": "iic/punc_ct-transformer_zh-cn-common-vocab272727-pytorch",
|
||||
"kind": "punctuation",
|
||||
"description": "Optional FunASR CT-Transformer punctuation restoration",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"config.yaml"
|
||||
],
|
||||
"any_files": [
|
||||
"*.pt",
|
||||
"*.bin",
|
||||
"*.onnx"
|
||||
],
|
||||
"required_files": ["configuration.json", "config.yaml", "model.pb"],
|
||||
"min_total_size_bytes": 1000000
|
||||
},
|
||||
"iic/speech_campplus_speaker-diarization_common": {
|
||||
|
|
@ -82,11 +43,7 @@
|
|||
"kind": "speaker_verification",
|
||||
"description": "Configured CAM++ speaker verification",
|
||||
"revision": "v2.0.2",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"config.yaml",
|
||||
"campplus_cn_common.bin"
|
||||
],
|
||||
"required_files": ["configuration.json", "config.yaml", "campplus_cn_common.bin"],
|
||||
"min_total_size_bytes": 10000000
|
||||
},
|
||||
"iic/speech_eres2netv2_sv_zh-cn_16k-common": {
|
||||
|
|
@ -94,12 +51,8 @@
|
|||
"directory": "iic/speech_eres2netv2_sv_zh-cn_16k-common",
|
||||
"kind": "realtime_speaker_verification",
|
||||
"description": "Realtime speaker verification",
|
||||
"required_files": [
|
||||
"configuration.json"
|
||||
],
|
||||
"any_files": [
|
||||
"*"
|
||||
],
|
||||
"required_files": ["configuration.json"],
|
||||
"any_files": ["*"],
|
||||
"min_total_size_bytes": 10000000
|
||||
},
|
||||
"damo/speech_campplus_sv_zh-cn_16k-common": {
|
||||
|
|
@ -107,11 +60,7 @@
|
|||
"directory": "damo/speech_campplus_sv_zh-cn_16k-common",
|
||||
"kind": "speaker_verification",
|
||||
"description": "CAM++ speaker verification dependency",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"config.yaml",
|
||||
"campplus_cn_common.bin"
|
||||
],
|
||||
"required_files": ["configuration.json", "config.yaml", "campplus_cn_common.bin"],
|
||||
"min_total_size_bytes": 10000000
|
||||
},
|
||||
"damo/speech_campplus-transformer_scl_zh-cn_16k-common": {
|
||||
|
|
@ -119,11 +68,7 @@
|
|||
"directory": "damo/speech_campplus-transformer_scl_zh-cn_16k-common",
|
||||
"kind": "speaker_transformer",
|
||||
"description": "CAM++ Transformer dependency",
|
||||
"required_files": [
|
||||
"configuration.json",
|
||||
"campplus_cn_encoder.pt",
|
||||
"transformer_backend.pt"
|
||||
],
|
||||
"required_files": ["configuration.json", "campplus_cn_encoder.pt", "transformer_backend.pt"],
|
||||
"min_total_size_bytes": 10000000
|
||||
},
|
||||
"Qwen/Qwen3-ForcedAligner-0.6B": {
|
||||
|
|
@ -131,13 +76,8 @@
|
|||
"directory": "Qwen/Qwen3-ForcedAligner-0.6B",
|
||||
"kind": "forced_aligner",
|
||||
"description": "Qwen3 word-level forced aligner",
|
||||
"required_files": [
|
||||
"config.json"
|
||||
],
|
||||
"any_files": [
|
||||
"*.safetensors",
|
||||
"*.bin"
|
||||
],
|
||||
"required_files": ["config.json"],
|
||||
"any_files": ["*.safetensors", "*.bin"],
|
||||
"min_total_size_bytes": 500000000
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,23 +3,18 @@ requires = ["setuptools>=68", "wheel"]
|
|||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "funasr-realtime-asr-demo"
|
||||
name = "qwen3-asr-vllm-deployment"
|
||||
version = "0.1.0"
|
||||
description = "FunASR streaming ASR browser demo"
|
||||
description = "Standalone Qwen3-ASR model downloader and VLLM deployment"
|
||||
requires-python = ">=3.10,<3.14"
|
||||
dependencies = [
|
||||
"aiohttp==3.11.11",
|
||||
"python-dotenv>=1.0",
|
||||
"funasr==1.4.16",
|
||||
"modelscope[framework]==1.34.0",
|
||||
"soundfile==0.13.1",
|
||||
"librosa==0.11.0",
|
||||
"numpy>=1.24",
|
||||
"modelscope==1.34.0",
|
||||
"qwen-asr[vllm]==0.0.6",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
funasr-realtime-demo = "backend.run_funasr_demo:main"
|
||||
funasr-download-models = "scripts.download_models:main"
|
||||
qwen3-asr-download = "scripts.download_models:main"
|
||||
qwen3-asr-serve = "scripts.serve:main"
|
||||
|
||||
[tool.setuptools]
|
||||
packages = ["scripts", "backend", "backend.realtime_websocket"]
|
||||
packages = ["scripts"]
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -0,0 +1,87 @@
|
|||
# Demo 说话人链路修复与验收
|
||||
|
||||
本次仅修改 `demo/`。参考原 `app/services/qwen3_websocket_asr.py` 的前滚音频、同句异步说话人更新和结束提交方式,保留独立 vLLM ASR + 辅助声纹服务 + WebSocket 三个进程。当前机器没有显卡,验证采用模拟模型,真实声纹准确率和模型加载状态需在部署机器上验收。
|
||||
|
||||
## 已确认的代码问题
|
||||
|
||||
1. 页面只处理 `sentences`,忽略已实现的 `display_state`,因此后端合并结果没有展示。未知片段还会临时附在上一位已确认说话人的气泡里。
|
||||
2. 页面点击停止后 5 秒强制断开,但辅助 HTTP 请求的超时是 45 秒,可能截掉迟到的声纹更新。
|
||||
3. 无 embedding、低置信度、协议字段不完整等情况被静默丢弃,用户无法区分等待、短音频、服务错误和证据拒绝。
|
||||
4. 旧声纹长度检查把句尾 800ms 静音一起计算,可能让短插话通过长度要求;WAV 固定去掉 44 字节也会破坏带额外元数据的输入。
|
||||
5. 仅 vLLM 启动器加载 `demo/.env`,另两个进程没有加载;`0.6b` 下载别名也可能被直接当作公开模型名发送。
|
||||
|
||||
这些是代码中可复现的问题,不能据此断言部署中的声纹模型一定已经正常加载。新版本将具体原因直接显示在片段气泡和日志里。
|
||||
|
||||
## 修复后的行为
|
||||
|
||||
- partial/final 和声纹更新继续使用同一 `sentence_id`。页面显示完整聚合快照,过时的 `revision` 不覆盖新快照。
|
||||
- 仅合并相邻且身份可信的片段。未知短插话独立显示;A→B→A 保留顺序。不同实名或实名与弱匿名身份不会因为簇 ID 相同而合并。
|
||||
- 每条声纹证据必须来自当前片段;同步、异步两种入口都拒绝 `short_attach` / `embedding_attach`。不复制上一段 embedding。声纹向量拒绝零向量、NaN、Infinity 和多样本矩阵。
|
||||
- WebSocket 以有效有声帧判断长度。少于 800ms 的语音保持 pending;800ms~1.6s 也必须独立提取特征,不能直接继承前一位身份。该长度门槛属于保守保护,不能保证短样本识别准确率。
|
||||
- 保留 200ms 前滚,提交时去除尾部静音。增量解析 WAV 的 RIFF/fmt/data/JUNK 等头,拒绝非 16kHz 单声道 PCM16 文件。
|
||||
- `stop/eof` 返回 `draining`,排空 ASR 和声纹队列后才发送最终快照及 `end`。声纹失败不阻断 ASR。abort、断线、异常和正常结束均清理聚类会话。
|
||||
- ASR worker 异常会立即发送 `error`,不会一直等待客户端停止。
|
||||
|
||||
## 如何判断卡在哪里
|
||||
|
||||
页面顶部显示已确认数量;未匹配到说话人的气泡统一显示“未知说话人”,详细等待或失败原因通过标签悬停提示和原始日志查看。原始日志保留 `speaker_status`、`speaker_reason`、`speaker_strategy` 和置信度。
|
||||
|
||||
| `speaker_status` | 含义与排查方向 |
|
||||
|---|---|
|
||||
| `waiting_final` | 讲话仍在进行,等待静音或最大时长切段 |
|
||||
| `queued` / `processing` | 文本已完成,声纹正在排队或推理 |
|
||||
| `confirmed` | 当前片段声纹已确认,应该显示说话人标签 |
|
||||
| `insufficient_audio` | 有效语音过短,不继承上一位;使用较长发言复测 |
|
||||
| `service_unavailable` / `service_error` | 未配置、无法访问或模型推理失败;检查辅助服务及完整错误 |
|
||||
| `no_embedding` | 服务没有产生可用特征 |
|
||||
| `evidence_rejected` | 缺少 fresh/confirmed 证据、置信度不足或使用了继承策略 |
|
||||
| `disabled` | 本次未开启说话人分离 |
|
||||
|
||||
健康检查为 `http://辅助服务器:8010/health`。新版本有 `speaker_protocol_version: 2`,重点检查 `speaker_embedding_ready` 与 `speaker_embedding_model`。若没有版本字段,检查是否重启了更新后的辅助进程。HTTP 健康检查不执行真实声纹推理,不能代替音频验收。
|
||||
|
||||
辅助服务目前返回匿名的“说话人 1、2……”;demo 没有接入原应用的声纹注册库,因此不会自动识别人员实名。vLLM 仅输出转写文本,声纹标签由 `scripts/auxiliary_server.py` 负责。
|
||||
|
||||
模型职责要区分:`iic/speech_campplus_sv_zh-cn_16k-common` 是实时 turn 的 CAM++ embedding 模型,必须加载;`iic/speech_campplus_speaker-diarization_common` 是完整音频分离 pipeline,包含额外的 change locator/VAD 依赖,当前实时 WebSocket 不在启动阶段调用它。WebSocket 自身仍用轻量 RMS 帧门控切句,辅助服务的 FunASR VAD 对外提供 `/v1/vad`,并供完整 diarization 依赖使用;因此启动辅助服务是为了 CAM++ 声纹和聚类,不能把整段 diarization 的加载失败误认为 vLLM 失败。
|
||||
|
||||
## 部署后操作
|
||||
|
||||
在部署机更新这些文件后,已有 vLLM 可继续运行。重新安装增补的 `python-dotenv` 依赖,并重启辅助服务及 WebSocket。辅助服务启动时强制依赖 VAD + CAM++ `speaker_verification`;完整 CAM++ diarization 不再阻断实时启动;以下命令都从 `demo/` 目录执行:
|
||||
|
||||
```text
|
||||
pip install -r requirements-auxiliary.txt
|
||||
pip install -r realtime_asr_optimization_demo/requirements.txt
|
||||
python scripts/auxiliary_server.py
|
||||
```
|
||||
|
||||
如果仍提示核心模型缺失或加载失败,日志会列出模型 ID、实际查找路径、状态和底层异常。先执行 `python scripts/download_models.py --auxiliary-only`,或设置 `.env` 的 `MODEL_DIR` 指向同时包含 `damo/speech_fsmn_vad_zh-cn-16k-common-pytorch` 与 `iic/speech_campplus_sv_zh-cn_16k-common` 的目录。不要用 `AUXILIARY_ALLOW_MISSING=true` 掩盖 VAD/CAM++ 核心模型缺失;该选项只适合临时查看可选模型状态。
|
||||
|
||||
另一个终端执行:
|
||||
|
||||
```text
|
||||
python realtime_asr_optimization_demo/server.py --no-browser
|
||||
```
|
||||
|
||||
三个进程现在均读取 `demo/.env`,系统环境变量优先。确认 `MODEL_SERVICE_URL` 指向现有 vLLM 的 `/v1`,`AUXILIARY_SERVICE_URL` 指向辅助服务。`QWEN3_ASR_MODEL` 支持部署清单别名,设置 `VLLM_SERVED_MODEL_NAME` 时优先使用该公开名称。更新后刷新浏览器;脚本 URL 已更新版本号。
|
||||
|
||||
`127.0.0.1` 指各 Python 服务运行的机器,不是浏览器所在机器。跨服务器部署时填写对应服务器 IP。浏览器麦克风访问远程页面需要安全上下文(HTTPS);本机 localhost 可用于测试。
|
||||
|
||||
## 验收顺序
|
||||
|
||||
1. 单人讲话 2~4 秒后停顿:先出现文本,随后同句变为“说话人 1”。
|
||||
2. 同一人再次讲话并停顿:确认后相邻块应合并;取消页面“合并相邻”可对照物理片段。
|
||||
3. A→B→A,各说 2 秒以上:应保持三个时间顺序块。标签准确率需要真实声纹模型验证。
|
||||
4. A 后 B 说一个很短的“嗯”:应显示“有效语音不足”,不能进入 A 的气泡。
|
||||
5. 辅助服务关闭时测试:ASR 仍完成,片段明确显示服务错误。
|
||||
6. 讲话中点击停止:等待最终结果,不能在五秒时丢失说话人更新。
|
||||
|
||||
无显卡回归命令:
|
||||
|
||||
```text
|
||||
cd demo
|
||||
python -m unittest discover -s tests -v
|
||||
cd realtime_asr_optimization_demo
|
||||
python -m unittest discover -s tests -v
|
||||
node --test tests/test_frontend.cjs
|
||||
```
|
||||
|
||||
这些测试覆盖状态机、模拟 HTTP/WebSocket、延迟更新、前端脚本和有效向量校验,不执行模型下载或 GPU 推理。当前 VLLM 适配器默认使用 `/v1/realtime` 原生流式,旧同步 HTTP 保留为兼容回退。单个内部片段中无停顿的多人换话或重叠讲话仍需真实模型和更细粒度切段验证。
|
||||
|
|
@ -0,0 +1,135 @@
|
|||
# Realtime ASR WebSocket Optimization Demo
|
||||
|
||||
这是一个独立的实时 ASR WebSocket 验证项目,放在模型部署项目 `demo` 下,但不导入、不启动、也不调用仓库根目录的原始 `app/`。
|
||||
|
||||
## 验证目标
|
||||
|
||||
- 浏览器麦克风或按实时速度发送的 PCM/WAV 音频通过本项目 WebSocket 发送。
|
||||
- 同一 `sentence_id` 的 partial、final 和后续更新覆盖同一条 raw segment。
|
||||
- raw segment 与前端 display block 分离,避免把物理切段直接等同于展示换行。
|
||||
- 小于 1.6 秒且带有 `short_attach` / `embedding_attach` 策略的实名结果降级为 pending。
|
||||
- 不把 embedding 字段写入 demo 状态池。
|
||||
- 只有相邻且身份可信的 segment 才合并;A→B→A 保持时间顺序。
|
||||
- 记录 partial 首次延迟、final 延迟、partial 修订次数和服务端返回时间范围。
|
||||
- 每个实时 turn 单独提交声纹特征,由辅助服务维护本 WebSocket session 的在线聚类中心。
|
||||
|
||||
本次修复、诊断状态和部署验收步骤见 [FIXES.md](FIXES.md)。ASR 继续使用已部署的独立 vLLM;已有 vLLM 服务时无需重复启动或下载模型。
|
||||
|
||||
当前默认使用 vLLM `/v1/realtime` WebSocket:PCM 音频帧只发送一次,模型
|
||||
持续返回 `transcription.delta`,VAD 切段时发送 `input_audio_buffer.commit`。
|
||||
因此 partial 不再重复上传不断增长的整段音频,适合长时间会议运行。
|
||||
启动脚本默认启用 `Qwen3ASRRealtimeGeneration`;如果部署端点不支持 realtime,
|
||||
本项目会自动回退到原同步 HTTP 接口,外部 WebSocket 参数不变。
|
||||
|
||||
## 启动
|
||||
|
||||
需要启动两个模型服务和一个 WebSocket 页面:VLLM 只负责 ASR,辅助服务负责
|
||||
VAD、CAM++ 声纹模型及在线聚类,WebSocket 只做音频流编排,不导入原项目代码。
|
||||
|
||||
先在 `demo` 目录下载 ASR 和辅助模型,并启动 VLLM:
|
||||
|
||||
```powershell
|
||||
cd D:\github-project\ASR\Qwen-Asr\demo
|
||||
python scripts\download_models.py
|
||||
python scripts\serve.py
|
||||
```
|
||||
|
||||
另一个终端启动辅助模型服务:
|
||||
|
||||
```powershell
|
||||
cd D:\github-project\ASR\Qwen-Asr\demo
|
||||
pip install -r requirements-auxiliary.txt
|
||||
python scripts\auxiliary_server.py
|
||||
```
|
||||
|
||||
另开一个终端启动 WebSocket 页面:
|
||||
|
||||
```powershell
|
||||
cd D:\github-project\ASR\Qwen-Asr\demo\realtime_asr_optimization_demo
|
||||
python -m venv .venv
|
||||
.\.venv\Scripts\Activate.ps1
|
||||
pip install -r requirements.txt
|
||||
python server.py --no-browser
|
||||
```
|
||||
|
||||
页面服务默认监听 `0.0.0.0:8082`,端口在 `server.py` 顶部的 `WEB_PORT` 内部变量中维护。
|
||||
VLLM 默认地址为 `http://127.0.0.1:9950/v1`,辅助服务默认地址为
|
||||
`http://127.0.0.1:8010`。可通过环境变量切换到远程服务:
|
||||
|
||||
```powershell
|
||||
$env:MODEL_SERVICE_URL = 'http://127.0.0.1:9950/v1'
|
||||
$env:AUXILIARY_SERVICE_URL = 'http://127.0.0.1:8010'
|
||||
python server.py --no-browser
|
||||
```
|
||||
|
||||
如果服务器端口已通过 VS Code Remote/端口转发映射到本机,保持上述两个
|
||||
`127.0.0.1` 地址即可:本地 WebSocket 只负责编排,ASR 和 VAD/CAM++ 推理仍在
|
||||
服务器 GPU 服务中完成。先访问 `http://127.0.0.1:9950/v1/models` 与
|
||||
`http://127.0.0.1:8010/health`,分别确认 VLLM 模型和辅助模型服务可达且 `ready=true`。
|
||||
|
||||
服务器部署时使用 `--no-browser`,然后在客户端浏览器访问 `http://服务器IP:8082`。如需让启动日志显示服务器域名或 IP,可设置 `WEB_DISPLAY_HOST`;它只影响提示文本,不改变监听地址。
|
||||
|
||||
辅助服务启动后可用 `http://服务器IP:8010/health` 检查模型状态。WebSocket
|
||||
收到聚类服务错误时仍会继续输出 ASR,但对应片段会显示“未知说话人”;详细的
|
||||
`speaker_reason` 可将鼠标悬停在标签上查看,事件日志仍会显示 `speaker_warning`,
|
||||
便于区分“模型未归类”和“ASR 失败”。
|
||||
|
||||
声纹服务的实时路径必须通过 ModelScope pipeline 的公开接口提取 embedding:
|
||||
`pipeline([wav_path], output_emb=True)`。不能绕过 pipeline 预处理后直接调用
|
||||
`pipeline.model`,否则采样率、声道和 waveform 预处理不会执行,部分 ModelScope
|
||||
版本会直接抛异常,WebSocket 仍会继续输出 ASR 并把说话人保留为 pending。
|
||||
更新辅助服务代码后需要重启 `python scripts/auxiliary_server.py`,仅重启页面
|
||||
服务不会替换已经驻留在 GPU 中的旧辅助服务进程。
|
||||
|
||||
## 页面操作
|
||||
|
||||
1. 选择 Mic 或 File。
|
||||
2. 点击开始,浏览器通过 `/ws` 建立本项目 WebSocket。
|
||||
3. 页面展示腾讯 Demo 风格的气泡;麦克风或 PCM/WAV 文件按流式方式输入,partial 会在讲话过程中实时刷新。
|
||||
4. VAD 检测到静音后提交当前 turn,先返回 pending,再异步更新说话人。
|
||||
5. 停止会发送 `stop`,服务端完成当前 turn 和 speaker 队列后再发送 `end`。
|
||||
6. `abort` 只取消会话,不提交当前片段。
|
||||
|
||||
## WebSocket 消息
|
||||
|
||||
客户端首条消息:
|
||||
|
||||
```json
|
||||
{
|
||||
"type": "start",
|
||||
"source": "mic",
|
||||
"model_service_url": "http://127.0.0.1:9950/v1",
|
||||
"model": "Qwen/Qwen3-ASR-0.6B",
|
||||
"speaker_diarization": 1,
|
||||
"sentence_strategy": 0,
|
||||
"enable_native_partial_stream": true,
|
||||
"partial_interval_ms": 1200,
|
||||
"max_segment_sec": 12,
|
||||
"speaker_gap_ms": 400,
|
||||
"display_merge": true
|
||||
}
|
||||
```
|
||||
|
||||
随后持续发送 16kHz、单声道、PCM16 二进制音频;文件模式仅支持 PCM/WAV,结束发送 `{"type":"eof"}`,停止发送
|
||||
`{"type":"stop"}`,取消发送 `{"type":"abort"}`。每个已完成 turn 会向辅助服务
|
||||
发送一次 `/v1/speaker/resolve`,只包含当前 turn 音频和 session_id,不会重复上传整段会话。
|
||||
|
||||
服务端会发送 `start`、`sentences`、`display_state`、`metrics`、`speaker_warning`、`draining`、`end` 和 `error`。
|
||||
页面用带 `revision` 的 `display_state` 渲染,以 `block_id` 标识展示块;`sentences` 保留原始片段及诊断状态。
|
||||
停止后必须等待 `end`,其中包含完整 `sentences` 和 `display_blocks`,不能提前关闭连接。
|
||||
`sentences` 中 `sentence_type=0` 是 partial,`sentence_type=1` 是 final;同一个
|
||||
`sentence_id` 必须覆盖更新而不是追加。`sentence_strategy=0` 使用约 800ms
|
||||
静音切句,`sentence_strategy=1` 使用约 1400ms 静音切段,更适合段落模式。
|
||||
|
||||
说话人模式下,`speaker_gap_ms`(默认 400ms)会在有效语音达到 800ms 后
|
||||
提前结束短停顿话轮,用于拆开 A→B 的交接;它不触发额外声纹请求,同一说话人的
|
||||
相邻片段仍由辅助服务聚类为同一说话人。服务未就绪时该策略会自动关闭。
|
||||
|
||||
## 目录
|
||||
|
||||
- `server.py`:本地 HTTP 页面和 WebSocket 会话编排。
|
||||
- `model_service.py`:独立 VLLM OpenAI 音频接口适配层。
|
||||
- `auxiliary_service.py`:独立辅助模型 HTTP 接口适配层。
|
||||
- `speaker_assembler.py`:raw segment、speaker evidence 和 display block 状态机。
|
||||
- `static/`:麦克风/文件测试页面。
|
||||
- `tests/`:只测试本项目状态合并和音频转换逻辑。
|
||||
|
|
@ -64,23 +64,6 @@ class AuxiliaryModelService:
|
|||
raise RuntimeError("auxiliary health check returned a non-object JSON value")
|
||||
return decoded
|
||||
|
||||
async def punctuate(self, text: str) -> dict[str, Any]:
|
||||
"""Apply optional FunASR punctuation without making it a startup dependency."""
|
||||
if self._session is None:
|
||||
raise RuntimeError("auxiliary model service is not started")
|
||||
endpoint = self.config.base_url.rstrip("/") + "/v1/punctuation"
|
||||
async with self._session.post(endpoint, json={"text": text}) as response:
|
||||
body = await response.text()
|
||||
if response.status >= 400:
|
||||
raise RuntimeError(f"auxiliary punctuation failed ({response.status}): {body[:500]}")
|
||||
try:
|
||||
decoded = await response.json(content_type=None)
|
||||
except ValueError as exc:
|
||||
raise RuntimeError(f"auxiliary punctuation returned invalid JSON: {body[:500]}") from exc
|
||||
if not isinstance(decoded, dict):
|
||||
raise RuntimeError("auxiliary punctuation returned a non-object JSON value")
|
||||
return decoded
|
||||
|
||||
async def resolve_speaker(
|
||||
self,
|
||||
pcm_bytes: bytes,
|
||||
|
|
@ -116,35 +99,6 @@ class AuxiliaryModelService:
|
|||
# 保留无标签响应里的具体原因;由组装器统一判断可信度,避免这里静默丢弃。
|
||||
return decoded
|
||||
|
||||
async def track_speakers(
|
||||
self,
|
||||
pcm_bytes: bytes,
|
||||
session_id: str,
|
||||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""对已完成的 FunASR turn 按 CAM++ 滑窗追踪段内说话人。"""
|
||||
if self._session is None:
|
||||
raise RuntimeError("auxiliary model service is not started")
|
||||
form = FormData()
|
||||
form.add_field("file", pcm16_to_wav(pcm_bytes), filename="turn.wav", content_type="audio/wav")
|
||||
form.add_field("session_id", session_id)
|
||||
form.add_field("start_time_ms", str(start_time_ms))
|
||||
form.add_field("end_time_ms", str(end_time_ms))
|
||||
endpoint = self.config.base_url.rstrip("/") + "/v1/speaker/track"
|
||||
async with self._session.post(endpoint, data=form) as response:
|
||||
body = await response.text()
|
||||
if response.status >= 400:
|
||||
raise RuntimeError(f"auxiliary speaker tracking failed ({response.status}): {body[:500]}")
|
||||
try:
|
||||
decoded = await response.json(content_type=None)
|
||||
except ValueError as exc:
|
||||
raise RuntimeError(f"auxiliary speaker tracking returned invalid JSON: {body[:500]}") from exc
|
||||
segments = decoded.get("segments", []) if isinstance(decoded, dict) else []
|
||||
if not isinstance(segments, list):
|
||||
return []
|
||||
return [segment for segment in segments if isinstance(segment, dict)]
|
||||
|
||||
async def reset_speaker_session(self, session_id: str) -> None:
|
||||
"""通知辅助服务释放当前 WebSocket 对应的在线聚类状态。"""
|
||||
if self._session is None:
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
aiohttp==3.11.11
|
||||
python-dotenv>=1.0
|
||||
|
|
@ -25,7 +25,7 @@ from speaker_assembler import SegmentAssembler
|
|||
|
||||
|
||||
# 与部署启动器读取同一配置;外部环境变量优先于 demo/.env。
|
||||
DEPLOY_ROOT = Path(__file__).resolve().parents[2]
|
||||
DEPLOY_ROOT = Path(__file__).resolve().parents[1]
|
||||
load_dotenv(DEPLOY_ROOT / ".env")
|
||||
|
||||
# 监听所有网卡,允许同一局域网内的浏览器访问服务器上的 Demo;端口集中在代码
|
||||
|
|
@ -1,5 +1,8 @@
|
|||
// ===== DOM Elements =====
|
||||
// ===== 页面元素 =====
|
||||
const elEngineModel = document.getElementById('engineModel');
|
||||
const elModelServiceUrl = document.getElementById('modelServiceUrl');
|
||||
const elSpeakerStatus = document.getElementById('speakerStatus');
|
||||
const elDisplayMerge = document.getElementById('displayMerge');
|
||||
const elSpeakerDiarization = document.getElementById('speakerDiarization');
|
||||
const elDiarizationLabel = document.getElementById('diarizationLabel');
|
||||
const elSentenceStrategy = document.getElementById('sentenceStrategy');
|
||||
|
|
@ -19,13 +22,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 +39,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) {
|
||||
|
|
@ -52,12 +55,20 @@ function appendLog(msg) {
|
|||
const typeClass = 'log-type-' + (msg.type || 'unknown');
|
||||
const entry = document.createElement('div');
|
||||
entry.className = 'log-entry';
|
||||
entry.innerHTML = `<span class="log-time">${ts}</span><span class="${typeClass}">${JSON.stringify(msg)}</span>`;
|
||||
// 原始文本不作为 HTML 解释,转写中的标签也应原样显示。
|
||||
const stamp = document.createElement('span');
|
||||
stamp.className = 'log-time';
|
||||
stamp.textContent = ts;
|
||||
const content = document.createElement('span');
|
||||
content.className = typeClass;
|
||||
content.textContent = JSON.stringify(msg);
|
||||
entry.append(stamp, content);
|
||||
elLogArea.appendChild(entry);
|
||||
while (elLogArea.childNodes.length > 300) elLogArea.firstChild.remove();
|
||||
elLogArea.scrollTop = elLogArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== State =====
|
||||
// ===== 会话状态 =====
|
||||
let ws = null;
|
||||
let sending = false;
|
||||
let stoppingByUser = false;
|
||||
|
|
@ -71,7 +82,7 @@ let micWorklet = null;
|
|||
let micTimerInterval = null;
|
||||
let micStartTime = 0;
|
||||
|
||||
// Input mode (mic / file)
|
||||
// 输入模式(麦克风 / 文件)
|
||||
let inputMode = 'mic';
|
||||
let selectedFile = null;
|
||||
|
||||
|
|
@ -80,23 +91,28 @@ const EXT_FORMAT_MAP = {
|
|||
'pcm': 1, 'wav': 12, 'mp3': 8, 'm4a': 14,
|
||||
'aac': 16, 'opus': 10, 'ogg': 10, 'silk': 6, 'speex': 4
|
||||
};
|
||||
// 不同格式的默认发送倍速:PCM/WAV 实时速度 1x,压缩格式解压快可提速
|
||||
// PCM/WAV 的默认发送倍速;实时验证默认按 1 倍速输入。
|
||||
const DEFAULT_SPEED = {
|
||||
'pcm': 1.0, 'wav': 1.0,
|
||||
'mp3': 2.0, 'm4a': 2.0, 'aac': 2.0,
|
||||
'opus': 3.0, 'ogg': 3.0, 'silk': 3.0, 'speex': 3.0
|
||||
};
|
||||
const MAX_SPEED = 3.0;
|
||||
const UNSUPPORTED_STREAMING_EXTENSIONS = new Set(['m4a']);
|
||||
// 实时 WebSocket 需要服务端逐帧读取音频;压缩格式必须等文件完整后才能解码,
|
||||
// 因此本次流式验证只允许 PCM/WAV,避免把整段上传伪装成实时识别。
|
||||
const STREAMABLE_AUDIO_EXTENSIONS = new Set(['pcm', 'wav']);
|
||||
|
||||
const SPEAKER_COLORS = ['#4a7dff', '#52c41a', '#faad14', '#ff4d4f', '#9254de', '#13c2c2'];
|
||||
|
||||
let sentenceMap = {};
|
||||
let speakerOrderMap = {};
|
||||
let speakerOrderCounter = 0;
|
||||
let displayStateSupported = false;
|
||||
let displayRevision = -1;
|
||||
|
||||
// ===== Input Mode Tabs =====
|
||||
// ===== 输入模式标签页 =====
|
||||
function switchMode(mode) {
|
||||
if (ws) return;
|
||||
inputMode = mode;
|
||||
elTabMic.classList.toggle('active', mode === 'mic');
|
||||
elTabFile.classList.toggle('active', mode === 'file');
|
||||
|
|
@ -111,20 +127,20 @@ 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;
|
||||
const ext = getFileExt(file.name);
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
if (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
selectedFile = null;
|
||||
e.target.value = '';
|
||||
elFileInfo.textContent = 'M4A 暂不支持直接上传,请先转成 WAV 或 MP3';
|
||||
elFileInfo.textContent = '实时测试只支持 PCM 或 WAV,请先转换音频格式';
|
||||
elFileInfo.classList.remove('has-file');
|
||||
elAudioMeta.style.display = 'none';
|
||||
elSpeedControl.style.display = 'none';
|
||||
elBtnStart.disabled = true;
|
||||
showToast('M4A 容器格式无法按当前实时切片方式直接识别,请转成 WAV 或 MP3', true);
|
||||
showToast('压缩音频不能按当前实时 WebSocket 逐帧识别,请转成 PCM 或 WAV', true);
|
||||
return;
|
||||
}
|
||||
selectedFile = file;
|
||||
|
|
@ -135,12 +151,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 +221,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 +246,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 +292,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,10 +309,10 @@ function setStatus(state, text) {
|
|||
elStatusText.textContent = text;
|
||||
}
|
||||
|
||||
// ===== Render: Subtitle (no diarization) =====
|
||||
// ===== 渲染字幕(关闭说话人分离) =====
|
||||
// 每个 sentence_id 对应一个独立气泡:
|
||||
// - 中间态(sentence_type===0):实时更新该气泡文本,末尾加 "..." 表示未定
|
||||
// - 稳态(sentence_type===1):去掉 "..." 并定格,下一个 sentence_id 自动创建新气泡
|
||||
// - 中间态(sentence_type===0):实时更新该气泡文本,末尾加 "..." 表示未定。
|
||||
// - 稳态(sentence_type===1):去掉 "..." 并定格,下一个 sentence_id 自动创建新气泡。
|
||||
function renderSubtitle(sentence) {
|
||||
const id = 'subtitle-' + sentence.sentence_id;
|
||||
const isInterim = sentence.sentence_type === 0;
|
||||
|
|
@ -331,127 +347,85 @@ function renderSubtitle(sentence) {
|
|||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== Render: Speaker Bubble =====
|
||||
// ===== 渲染说话人气泡 =====
|
||||
// 未确认片段独立展示,不能临时塞进上一位说话人的气泡。
|
||||
function renderBubble(sentence) {
|
||||
const id = 'sent-' + sentence.sentence_id;
|
||||
const speakerId = sentence.speaker_id;
|
||||
const isUnknown = speakerId < 0;
|
||||
const speakerId = Number(sentence.speaker_id);
|
||||
const trusted = Number.isInteger(speakerId) && speakerId >= 0
|
||||
&& ['fresh', 'confirmed'].includes(sentence.speaker_evidence);
|
||||
const isInterim = sentence.sentence_type === 0;
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
let entry = sentenceMap[id];
|
||||
let insertBefore = null;
|
||||
if (entry && entry.speakerId !== speakerId) {
|
||||
insertBefore = entry.el.nextSibling;
|
||||
entry.el.remove();
|
||||
entry = null;
|
||||
delete sentenceMap[id];
|
||||
}
|
||||
|
||||
if (!entry) {
|
||||
if (!(speakerId in speakerOrderMap)) {
|
||||
speakerOrderMap[speakerId] = speakerOrderCounter++;
|
||||
}
|
||||
const orderIdx = speakerOrderMap[speakerId];
|
||||
const side = orderIdx % 2 === 0 ? 'left' : 'right';
|
||||
const colorIdx = orderIdx % SPEAKER_COLORS.length;
|
||||
|
||||
const el = document.createElement('div');
|
||||
el.id = id;
|
||||
el.className = `bubble-row speaker-${side} speaker-${colorIdx}`;
|
||||
const wrapper = document.createElement('div');
|
||||
wrapper.className = 'bubble-wrapper';
|
||||
const header = document.createElement('div');
|
||||
header.className = 'bubble-header';
|
||||
const badge = document.createElement('span');
|
||||
badge.className = `speaker-badge speaker-color-${colorIdx}`;
|
||||
const nameSpan = document.createElement('span');
|
||||
nameSpan.className = 'speaker-name';
|
||||
nameSpan.textContent = `说话人 ${speakerId}`;
|
||||
const timeSpan = document.createElement('span');
|
||||
timeSpan.className = 'bubble-time';
|
||||
header.appendChild(badge);
|
||||
header.appendChild(nameSpan);
|
||||
header.appendChild(timeSpan);
|
||||
const body = document.createElement('div');
|
||||
body.className = 'bubble-body';
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
el.appendChild(wrapper);
|
||||
|
||||
if (insertBefore) elResultArea.insertBefore(el, insertBefore);
|
||||
else elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: speakerId };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
|
||||
const el = entry.el;
|
||||
const timeSpan = el.querySelector('.bubble-time');
|
||||
const body = el.querySelector('.bubble-body');
|
||||
timeSpan.textContent = formatTimeRange(sentence.start_time, sentence.end_time);
|
||||
Array.from(body.childNodes).filter(n => n.nodeType === Node.TEXT_NODE).forEach(n => n.remove());
|
||||
const textNode = document.createTextNode(sentence.sentence + (isInterim ? ' ...' : ''));
|
||||
body.insertBefore(textNode, body.firstChild);
|
||||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
entry.speakerId = speakerId;
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
function renderFallbackPendingBubble(id, text, isInterim, sentence) {
|
||||
let entry = sentenceMap[id];
|
||||
if (!entry) {
|
||||
const el = document.createElement('div');
|
||||
el.id = id;
|
||||
el.className = 'bubble-row speaker-left speaker-unknown';
|
||||
const wrapper = document.createElement('div');
|
||||
wrapper.className = 'bubble-wrapper';
|
||||
const header = document.createElement('div');
|
||||
header.className = 'bubble-header';
|
||||
const badge = document.createElement('span');
|
||||
badge.className = 'speaker-badge speaker-color-unknown';
|
||||
const nameSpan = document.createElement('span');
|
||||
nameSpan.className = 'speaker-name';
|
||||
nameSpan.textContent = '说话人不确定';
|
||||
const timeSpan = document.createElement('span');
|
||||
timeSpan.className = 'bubble-time';
|
||||
header.appendChild(badge);
|
||||
header.appendChild(nameSpan);
|
||||
header.appendChild(timeSpan);
|
||||
for (const name of ['speaker-badge', 'speaker-name', 'bubble-time']) {
|
||||
const span = document.createElement('span');
|
||||
span.className = name;
|
||||
header.appendChild(span);
|
||||
}
|
||||
const body = document.createElement('div');
|
||||
body.className = 'bubble-body';
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
wrapper.append(header, body);
|
||||
el.appendChild(wrapper);
|
||||
elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: -1 };
|
||||
entry = { el };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
entry.el.querySelector('.bubble-time').textContent = formatTimeRange(
|
||||
sentence.start_time, sentence.end_time
|
||||
);
|
||||
const body = entry.el.querySelector('.bubble-body');
|
||||
body.textContent = text;
|
||||
if (trusted && !(speakerId in speakerOrderMap)) speakerOrderMap[speakerId] = speakerOrderCounter++;
|
||||
const order = trusted ? speakerOrderMap[speakerId] : 0;
|
||||
const color = order % SPEAKER_COLORS.length;
|
||||
const el = entry.el;
|
||||
el.className = trusted ? `bubble-row speaker-${order % 2 ? 'right' : 'left'} speaker-${color}`
|
||||
: 'bubble-row speaker-left speaker-unknown';
|
||||
el.querySelector('.speaker-badge').className = 'speaker-badge speaker-color-' + (trusted ? color : 'unknown');
|
||||
// 未获得当前片段的可靠声纹证据时,标题保持简短;详细原因放到悬停提示,
|
||||
// 这样不会把“有效语音不足……”等内部诊断信息挤进说话人名称区域。
|
||||
const speakerName = el.querySelector('.speaker-name');
|
||||
speakerName.textContent = trusted
|
||||
? (sentence.speaker_name || `说话人 ${speakerId + 1}`)
|
||||
: '未知说话人';
|
||||
speakerName.title = trusted ? '' : (sentence.speaker_reason || '未匹配到说话人');
|
||||
el.querySelector('.bubble-time').textContent = formatTimeRange(sentence.start_time, sentence.end_time);
|
||||
const body = el.querySelector('.bubble-body');
|
||||
body.textContent = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// 按完整快照重建相邻块;序号防止两个后台 worker 的旧快照覆盖新状态。
|
||||
function renderDisplayState(msg, useSpeaker) {
|
||||
if (msg.revision != null && msg.revision <= displayRevision) return;
|
||||
if (msg.revision != null) displayRevision = msg.revision;
|
||||
elResultArea.replaceChildren();
|
||||
sentenceMap = {};
|
||||
const raw = msg.raw_segments || msg.sentences || [];
|
||||
if (useSpeaker) {
|
||||
for (const block of msg.display_blocks || []) renderBubble({ ...block, sentence_id: block.block_id });
|
||||
const confirmed = raw.filter(s => s.speaker_status === 'confirmed').length;
|
||||
const failed = raw.filter(s => ['service_error', 'service_unavailable', 'no_embedding', 'evidence_rejected'].includes(s.speaker_status)).length;
|
||||
elSpeakerStatus.textContent = `说话人:已确认 ${confirmed} / ${raw.length} 段` + (failed ? `,${failed} 段未识别成功(原因见气泡及日志)` : '');
|
||||
} else {
|
||||
raw.forEach(renderSubtitle);
|
||||
elSpeakerStatus.textContent = '说话人分离已关闭';
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Start Recognition =====
|
||||
elBtnStart.addEventListener('click', () => {
|
||||
if (inputMode === 'file' && !selectedFile) return;
|
||||
startRecognition();
|
||||
});
|
||||
|
||||
async function startRecognition() {
|
||||
if (ws) return;
|
||||
if (inputMode === 'file' && selectedFile) {
|
||||
const ext = getFileExt(selectedFile.name);
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
showToast('当前 demo 不支持直接流式上传 M4A,请转成 WAV 或 MP3', true);
|
||||
if (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
showToast('实时流式测试只支持 PCM 或 WAV,请转换后再试', true);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
|
@ -461,6 +435,9 @@ async function startRecognition() {
|
|||
sentenceMap = {};
|
||||
speakerOrderMap = {};
|
||||
speakerOrderCounter = 0;
|
||||
displayStateSupported = false;
|
||||
displayRevision = -1;
|
||||
elSpeakerStatus.textContent = '正在检查说话人服务…';
|
||||
audioChunks = [];
|
||||
elBtnExportWav.disabled = true;
|
||||
elResultPlaceholder?.remove();
|
||||
|
|
@ -474,8 +451,9 @@ async function startRecognition() {
|
|||
|
||||
const currentSession = ++sessionId;
|
||||
const useSpeaker = elSpeakerDiarization.checked;
|
||||
let receivedTerminal = false;
|
||||
|
||||
// 构造 start 消息
|
||||
// 构造 WebSocket 首条 start 消息。
|
||||
let voiceFormat = 0, fileName = '', speedFactor = 0;
|
||||
if (inputMode === 'file') {
|
||||
const ext = getFileExt(selectedFile.name);
|
||||
|
|
@ -486,7 +464,9 @@ async function startRecognition() {
|
|||
|
||||
const startPayload = {
|
||||
type: 'start',
|
||||
engine_model_type: elEngineModel.value,
|
||||
model: elEngineModel.value,
|
||||
model_service_url: elModelServiceUrl.value,
|
||||
display_merge: elDisplayMerge.checked,
|
||||
speaker_diarization: useSpeaker ? 1 : 0,
|
||||
sentence_strategy: parseInt(elSentenceStrategy.value),
|
||||
source: inputMode,
|
||||
|
|
@ -502,16 +482,14 @@ async function startRecognition() {
|
|||
ws.onopen = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
ws.send(JSON.stringify(startPayload));
|
||||
if (inputMode === 'file') {
|
||||
sendAudioFile(selectedFile);
|
||||
} else {
|
||||
startMicCapture();
|
||||
}
|
||||
// 在连接尚未建立时点击停止,也要在 start 后补发停止信号。
|
||||
if (!sending) ws.send(JSON.stringify({ type: 'stop' }));
|
||||
};
|
||||
|
||||
ws.onmessage = (event) => {
|
||||
if (currentSession !== sessionId) return;
|
||||
const msg = JSON.parse(event.data);
|
||||
if (msg.type === 'end' || msg.type === 'error') receivedTerminal = true;
|
||||
if (msg.type !== 'sentences') {
|
||||
console.log('[ws] type=' + msg.type, msg);
|
||||
}
|
||||
|
|
@ -530,8 +508,7 @@ async function startRecognition() {
|
|||
ws.onclose = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
stopMicCapture();
|
||||
if (sending) setStatus('error', '连接意外断开');
|
||||
else if (stoppingByUser) setStatus('done', '已停止');
|
||||
if (!receivedTerminal) setStatus('error', '连接中断,最终识别结果可能尚未完成');
|
||||
if (audioChunks.length > 0) elBtnExportWav.disabled = false;
|
||||
ws = null;
|
||||
resetControls();
|
||||
|
|
@ -548,10 +525,28 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
break;
|
||||
|
||||
case 'start':
|
||||
displayStateSupported = Boolean(msg.display_state_supported);
|
||||
currentVoiceId = msg.session_id;
|
||||
elVoiceIdDisplay.textContent = currentVoiceId || '—';
|
||||
elSpeakerStatus.textContent = useSpeaker
|
||||
? `说话人服务:${msg.speaker_service_url || '未配置'};片段结束后提取声纹`
|
||||
: '说话人分离已关闭';
|
||||
if (!sending) break;
|
||||
setStatus('running', '识别中...');
|
||||
if (inputMode === 'file') sendAudioFile(selectedFile).catch(handleInputError);
|
||||
else startMicCapture().catch(handleInputError);
|
||||
break;
|
||||
|
||||
case 'display_state':
|
||||
renderDisplayState(msg, useSpeaker);
|
||||
break;
|
||||
|
||||
case 'draining':
|
||||
setStatus('running', msg.message || '等待最终识别结果…');
|
||||
break;
|
||||
|
||||
case 'sentences':
|
||||
if (displayStateSupported) break;
|
||||
if (msg.sentences) {
|
||||
msg.sentences.forEach(s => {
|
||||
if (useSpeaker) renderBubble(s);
|
||||
|
|
@ -560,7 +555,14 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
break;
|
||||
|
||||
case 'speaker_warning':
|
||||
// ASR 仍可继续输出,但必须让测试人员立即知道说话人链路没有生效。
|
||||
elSpeakerStatus.textContent = '说话人服务异常:' + msg.message;
|
||||
showToast('说话人服务异常,详见状态和片段原因', true);
|
||||
break;
|
||||
|
||||
case 'end':
|
||||
if (msg.display_blocks) renderDisplayState(msg, useSpeaker);
|
||||
setStatus('done', '识别完成');
|
||||
sending = false;
|
||||
if (audioChunks.length > 0) elBtnExportWav.disabled = false;
|
||||
|
|
@ -578,7 +580,7 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== Stop =====
|
||||
// ===== 停止识别 =====
|
||||
elBtnStop.addEventListener('click', () => stopRecognition());
|
||||
|
||||
function stopRecognition() {
|
||||
|
|
@ -588,26 +590,12 @@ function stopRecognition() {
|
|||
setStatus('running', '停止中...');
|
||||
elBtnStop.disabled = true;
|
||||
|
||||
if (currentVoiceId) {
|
||||
fetch(`/api/stop?voice_id=${encodeURIComponent(currentVoiceId)}`)
|
||||
.then(r => r.json())
|
||||
.catch(err => console.error('[stop] error:', err));
|
||||
}
|
||||
|
||||
if (ws && ws.readyState === WebSocket.OPEN) {
|
||||
try { ws.send(JSON.stringify({ type: 'stop' })); } catch (e) {}
|
||||
}
|
||||
|
||||
const stopSession = sessionId;
|
||||
setTimeout(() => {
|
||||
if (sessionId !== stopSession) return;
|
||||
if (ws) {
|
||||
try { ws.close(); } catch (e) {}
|
||||
ws = null;
|
||||
setStatus('done', '已停止(超时)');
|
||||
resetControls();
|
||||
}
|
||||
}, 5000);
|
||||
// 等待服务端排空 ASR/声纹队列后发送 end,不能用五秒计时器截断更新。
|
||||
|
||||
}
|
||||
|
||||
function resetControls() {
|
||||
|
|
@ -622,36 +610,53 @@ function resetControls() {
|
|||
elBtnStop.disabled = true;
|
||||
}
|
||||
|
||||
// ===== Send Audio File =====
|
||||
// 按 16KB 切片发送,后端会缓冲成 6400 字节块并按 speed_factor 限流
|
||||
// ===== 发送音频文件 =====
|
||||
// 按 16KB 切片发送,并按照音频实际时长等待,确保文件模式也是真实的
|
||||
// 实时输入,而不是瞬间上传完整文件后再由服务端批量切片。
|
||||
const UPLOAD_CHUNK_SIZE = 16000;
|
||||
async function sendAudioFile(file) {
|
||||
const ownerSession = sessionId;
|
||||
const buffer = await file.arrayBuffer();
|
||||
if (ownerSession !== sessionId || !sending) return;
|
||||
const totalBytes = buffer.byteLength;
|
||||
let offset = 0;
|
||||
const ext = getFileExt(file.name);
|
||||
const isPcm = (ext === 'pcm');
|
||||
while (offset < totalBytes && sending && ws && ws.readyState === WebSocket.OPEN) {
|
||||
let bytesPerSecond = 16000 * 2;
|
||||
if (ext === 'wav' && totalBytes >= 44) {
|
||||
const header = new DataView(buffer, 0, 44);
|
||||
const byteRate = header.getUint32(28, true);
|
||||
if (byteRate > 0) bytesPerSecond = byteRate;
|
||||
}
|
||||
const speedFactor = Math.max(parseFloat(elSpeedSlider.value) || 1.0, 0.1);
|
||||
while (ownerSession === sessionId && offset < totalBytes && sending && ws && ws.readyState === WebSocket.OPEN) {
|
||||
const end = Math.min(offset + UPLOAD_CHUNK_SIZE, totalBytes);
|
||||
const chunk = buffer.slice(offset, end);
|
||||
// 仅 PCM 数据可直接拼成 WAV 导出;压缩格式跳过
|
||||
// 仅 PCM 数据可直接拼成 WAV 导出;当前实时模式不会接收压缩格式。
|
||||
if (isPcm) audioChunks.push(chunk.slice(0));
|
||||
ws.send(chunk);
|
||||
offset = end;
|
||||
await new Promise(r => setTimeout(r, 0));
|
||||
const chunkDurationMs = (chunk.byteLength / bytesPerSecond) * 1000 / speedFactor;
|
||||
await new Promise(r => setTimeout(r, Math.max(0, Math.round(chunkDurationMs))));
|
||||
}
|
||||
if (ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
if (ownerSession === sessionId && ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
sending = false;
|
||||
setStatus('running', '音频已发送,等待最终结果…');
|
||||
ws.send(JSON.stringify({ type: 'eof' }));
|
||||
}
|
||||
}
|
||||
|
||||
// ===== Microphone Capture =====
|
||||
// ===== 麦克风采集 =====
|
||||
async function startMicCapture() {
|
||||
const ownerSession = sessionId;
|
||||
let stream;
|
||||
try {
|
||||
micStream = await navigator.mediaDevices.getUserMedia({
|
||||
stream = await navigator.mediaDevices.getUserMedia({
|
||||
audio: { sampleRate: 16000, channelCount: 1, echoCancellation: true, noiseSuppression: true }
|
||||
});
|
||||
} catch (err) {
|
||||
if (ownerSession !== sessionId) return;
|
||||
handleInputError(err);
|
||||
console.error('getUserMedia error:', err);
|
||||
setStatus('error', '无法获取麦克风权限');
|
||||
elMicStatus.textContent = '无法获取麦克风: ' + err.message;
|
||||
|
|
@ -659,6 +664,11 @@ async function startMicCapture() {
|
|||
return;
|
||||
}
|
||||
|
||||
if (ownerSession !== sessionId || !sending) {
|
||||
stream.getTracks().forEach(track => track.stop());
|
||||
return;
|
||||
}
|
||||
micStream = stream;
|
||||
micAudioContext = new (window.AudioContext || window.webkitAudioContext)({ sampleRate: 16000 });
|
||||
const source = micAudioContext.createMediaStreamSource(micStream);
|
||||
const processor = micAudioContext.createScriptProcessor(4096, 1, 1);
|
||||
|
|
@ -713,3 +723,21 @@ function stopMicCapture() {
|
|||
elMicStatus.textContent = '点击下方按钮开始录音';
|
||||
elMicElapsed.textContent = '00:00';
|
||||
}
|
||||
|
||||
// 展示实际部署端点及模型,避免沿用旧 SDK 的无效引擎配置。
|
||||
fetch('/api/config').then(response => response.json()).then(config => {
|
||||
if (!ws) {
|
||||
elEngineModel.value = config.model;
|
||||
elModelServiceUrl.value = config.model_service_url;
|
||||
elSpeakerStatus.textContent = '说话人辅助服务:' + (config.speaker_service_url || '未配置');
|
||||
}
|
||||
}).catch(error => { elSpeakerStatus.textContent = '读取服务配置失败:' + error.message; });
|
||||
|
||||
// 输入端失败必须释放空会话,避免用户再次开始时留下旧连接。
|
||||
function handleInputError(error) {
|
||||
sending = false;
|
||||
setStatus('error', error.message);
|
||||
stopMicCapture();
|
||||
if (ws) { ws.close(); ws = null; }
|
||||
resetControls();
|
||||
}
|
||||
|
|
@ -20,15 +20,21 @@
|
|||
<div class="form-stack">
|
||||
<div class="form-group">
|
||||
<label for="engineModel">引擎模型</label>
|
||||
<input type="text" id="engineModel" value="16k_zh_en_speaker" readonly>
|
||||
<input type="text" id="engineModel" value="Qwen/Qwen3-ASR-0.6B">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="modelServiceUrl">vLLM 服务地址</label>
|
||||
<input type="text" id="modelServiceUrl" value="http://127.0.0.1:9950/v1">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="sentenceStrategy">分句策略</label>
|
||||
<select id="sentenceStrategy">
|
||||
<option value="0" selected>语义单句</option>
|
||||
<option value="1">段落</option>
|
||||
<option value="0" selected>短停顿(800ms)</option>
|
||||
<option value="1">长停顿(1400ms)</option>
|
||||
</select>
|
||||
</div>
|
||||
<!-- 展示合并与内部音频切段分开配置,便于对照原始片段。 -->
|
||||
<label><input type="checkbox" id="displayMerge" checked> 合并相邻且已确认的同一说话人</label>
|
||||
<div class="form-row">
|
||||
<div class="form-group">
|
||||
<label>话者分离</label>
|
||||
|
|
@ -54,9 +60,9 @@
|
|||
<div class="input-mode-panel" id="panelFile" style="display:none">
|
||||
<div class="file-select">
|
||||
<label class="btn btn-outline" for="audioFile">选择音频文件</label>
|
||||
<input type="file" id="audioFile" accept=".pcm,.wav,.mp3,.aac,.opus,.ogg,.silk,.speex" hidden>
|
||||
<input type="file" id="audioFile" accept=".pcm,.wav" hidden>
|
||||
</div>
|
||||
<span class="file-info" id="fileInfo">支持 pcm/wav/mp3/aac/opus/silk/speex,m4a 请先转 wav/mp3</span>
|
||||
<span class="file-info" id="fileInfo">支持 16kHz、单声道、PCM16 的 PCM/WAV</span>
|
||||
<div class="audio-meta" id="audioMeta" style="display:none">
|
||||
<div class="audio-meta-row"><span class="meta-k">格式</span><span class="meta-v" id="metaFormat">—</span></div>
|
||||
<div class="audio-meta-row"><span class="meta-k">采样率</span><span class="meta-v" id="metaSampleRate">—</span></div>
|
||||
|
|
@ -79,6 +85,7 @@
|
|||
<main class="panel-right">
|
||||
<section class="card result-card">
|
||||
<h2>识别结果</h2>
|
||||
<p id="speakerStatus" role="status">正在读取服务配置…</p>
|
||||
<div class="result-meta" id="resultMeta" style="display:none">
|
||||
<div class="meta-item">
|
||||
<span class="meta-label">VoiceID:</span>
|
||||
|
|
@ -106,6 +113,7 @@
|
|||
</main>
|
||||
</div>
|
||||
|
||||
<script src="app.js?v=210"></script>
|
||||
<!-- 版本号变更用于刷新浏览器缓存,确保加载“未知说话人”标签逻辑。 -->
|
||||
<script src="app.js?v=213"></script>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -801,3 +801,5 @@ input[type="text"][readonly]:focus {
|
|||
.log-type-end { color: #f9e2af; }
|
||||
.log-type-error { color: #f38ba8; }
|
||||
|
||||
|
||||
|
||||
|
|
@ -4,10 +4,7 @@ from __future__ import annotations
|
|||
|
||||
import unittest
|
||||
|
||||
try:
|
||||
from backend.realtime_websocket.speaker_assembler import SegmentAssembler
|
||||
except ModuleNotFoundError:
|
||||
from speaker_assembler import SegmentAssembler
|
||||
from speaker_assembler import SegmentAssembler
|
||||
|
||||
|
||||
class SegmentAssemblerTests(unittest.TestCase):
|
||||
|
|
@ -0,0 +1,6 @@
|
|||
aiohttp==3.11.11
|
||||
python-dotenv>=1.0
|
||||
funasr==1.3.1
|
||||
modelscope[framework]==1.34.0
|
||||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
|
|
@ -0,0 +1,7 @@
|
|||
# 使用新版 VLLM 原生支持 Qwen3-ASR,避免 qwen-asr[vllm] 将 VLLM 锁定到 0.14.0。
|
||||
--extra-index-url https://download.pytorch.org/whl/cu130
|
||||
torch==2.13.0
|
||||
torchvision==0.28.0
|
||||
torchaudio==2.11.0
|
||||
vllm==0.28.0
|
||||
python-dotenv>=1.0
|
||||
|
|
@ -1,17 +0,0 @@
|
|||
# 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.
|
||||
aiohttp==3.11.11
|
||||
python-dotenv>=1.0
|
||||
numpy>=1.24
|
||||
funasr==1.4.16
|
||||
modelscope[framework]==1.34.0
|
||||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
websockets>=12,<14
|
||||
|
||||
--extra-index-url https://download.pytorch.org/whl/cu130
|
||||
torch==2.13.0
|
||||
torchvision==0.28.0
|
||||
torchaudio==2.11.0
|
||||
vllm==0.28.0
|
||||
|
|
@ -6,10 +6,8 @@ from __future__ import annotations
|
|||
import asyncio
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Mapping
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
|
@ -18,28 +16,26 @@ from aiohttp import web
|
|||
from aiohttp.web_request import FileField
|
||||
from dotenv import load_dotenv
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from backend.model_manifest import auxiliary_models, load_manifest, model_directory
|
||||
try:
|
||||
from .model_manifest import auxiliary_models, load_manifest, model_directory
|
||||
except ImportError:
|
||||
from model_manifest import auxiliary_models, load_manifest, model_directory
|
||||
|
||||
|
||||
# 将辅助服务端口固定在代码变量中,服务器启动时只需执行脚本,便于部署和排查。
|
||||
AUXILIARY_HOST = "0.0.0.0"
|
||||
AUXILIARY_PORT = 8010
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
# 独立启动辅助服务也必须读取部署配置,不能只在启动 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.
|
||||
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")))
|
||||
ONLINE_SPEAKER_MATCH_THRESHOLD = 0.68
|
||||
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.
|
||||
DEFAULT_PRELOAD_KINDS = {"speaker_verification"}
|
||||
# 实时 WebSocket 必须使用 VAD + CAM++ 声纹模型进行在线聚类。完整的
|
||||
# speech_campplus_speaker-diarization_common 是整段离线 diarization 接口,
|
||||
# 与实时每个 turn 的 CAM++ embedding 不是同一加载路径;它仍可按需加载。
|
||||
DEFAULT_PRELOAD_KINDS = {"vad", "speaker_verification"}
|
||||
|
||||
|
||||
def _coerce_finite_float(value: object) -> float | None:
|
||||
|
|
@ -99,35 +95,26 @@ class AuxiliaryRuntime:
|
|||
self.inference_lock = asyncio.Lock()
|
||||
# 每个 WebSocket session 独立维护聚类中心,避免不同浏览器会话互相污染。
|
||||
self.speaker_clusters: dict[str, list[dict[str, Any]]] = {}
|
||||
# Keep a bounded rolling window history per WebSocket session, like FunASR's tracker.
|
||||
self.speaker_history: dict[str, dict[str, Any]] = {}
|
||||
self.speaker_last_seen: dict[str, float] = {}
|
||||
self._speaker_cluster_backend: Any | None = None
|
||||
self._speaker_postprocess: Any | None = None
|
||||
|
||||
def _preload_kinds(self) -> set[str]:
|
||||
"""读取需要在启动时加载的模型类型,默认不加载完整 diarization。"""
|
||||
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.
|
||||
# 无论环境变量如何设置,VAD 和 CAM++ speaker_verification 都是核心
|
||||
# 依赖;额外类型只会增加预加载项,不能绕过核心模型校验。
|
||||
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:
|
||||
"""只加载当前运行接口需要的模型;依赖模型和对齐模型先保持本地资产就绪。"""
|
||||
kind = str(config.get("kind") or "")
|
||||
if kind in {"vad", "punctuation"}:
|
||||
if kind == "vad":
|
||||
from funasr import AutoModel
|
||||
|
||||
# Keep punctuation on CPU by default so CAM++ can retain the GPU.
|
||||
device = (
|
||||
os.getenv("FUNASR_PUNC_DEVICE", "cpu")
|
||||
if kind == "punctuation"
|
||||
else AUXILIARY_DEVICE
|
||||
)
|
||||
return AutoModel(
|
||||
model=str(path),
|
||||
device=device,
|
||||
device=AUXILIARY_DEVICE,
|
||||
disable_update=True,
|
||||
disable_pbar=True,
|
||||
disable_log=True,
|
||||
|
|
@ -149,12 +136,9 @@ class AuxiliaryRuntime:
|
|||
preload_kinds = self._preload_kinds()
|
||||
loaded_kinds: set[str] = set()
|
||||
for model_id, config in self.assets.items():
|
||||
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.
|
||||
path = Path(os.environ["CAM_MODEL_PATH"]).resolve()
|
||||
record: dict[str, Any] = {"path": str(path), "asset_ready": _asset_ready(path, config)}
|
||||
kind = str(config.get("kind") or "")
|
||||
if kind not in preload_kinds:
|
||||
record["state"] = "optional_not_preloaded" if record["asset_ready"] else "optional_missing"
|
||||
record["preload"] = False
|
||||
|
|
@ -182,7 +166,7 @@ class AuxiliaryRuntime:
|
|||
else:
|
||||
record["state"] = "load_error"
|
||||
record["error"] = "model loader returned no model"
|
||||
if kind == "speaker_verification":
|
||||
if kind in {"vad", "speaker_verification"}:
|
||||
failures.append(model_id)
|
||||
except Exception as exc:
|
||||
record["state"] = "load_error"
|
||||
|
|
@ -198,11 +182,15 @@ 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 和 CAM++ speaker_verification 都是核心依赖,缺失/加载异常时
|
||||
# 立即失败并列出路径与底层错误,避免页面一直显示“未确认”。
|
||||
required_failures = [
|
||||
model_id for model_id in failures
|
||||
if self.assets[model_id].get("kind") == "speaker_verification"
|
||||
if self.assets[model_id].get("kind") == "vad"
|
||||
or (
|
||||
self.assets[model_id].get("kind") == "speaker_verification"
|
||||
and self._speaker_embedding_model_id() is None
|
||||
)
|
||||
]
|
||||
if required_failures:
|
||||
details = "; ".join(
|
||||
|
|
@ -212,8 +200,8 @@ class AuxiliaryRuntime:
|
|||
)
|
||||
raise RuntimeError(
|
||||
"Auxiliary core model is missing or failed to load: " + details
|
||||
+ ". Run python scripts/download_models.py --funasr-runtime, "
|
||||
"or set MODEL_DIR/CAM_MODEL_PATH to the local CAM++ asset."
|
||||
+ ". Run `python scripts/download_models.py --auxiliary-only` "
|
||||
"or set MODEL_DIR to the directory containing the downloaded assets."
|
||||
)
|
||||
|
||||
def _find_model(self, kind: str) -> Any:
|
||||
|
|
@ -255,31 +243,10 @@ class AuxiliaryRuntime:
|
|||
|
||||
async def vad(self, audio_path: str) -> Any:
|
||||
"""使用临时音频文件执行一次串行化的 VAD 推理。"""
|
||||
# Realtime VAD lives in the WebSocket process; this HTTP endpoint loads on demand.
|
||||
try:
|
||||
model = self._find_model("vad")
|
||||
except RuntimeError:
|
||||
model = self._load_optional_kind("vad")
|
||||
async with self.inference_lock:
|
||||
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."""
|
||||
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:
|
||||
output = output[0]
|
||||
if isinstance(output, Mapping):
|
||||
result = output.get("text")
|
||||
else:
|
||||
result = getattr(output, "text", None)
|
||||
return str(result).strip() if result else text
|
||||
|
||||
async def diarization(self, audio_path: str) -> Any:
|
||||
"""使用 CAM++ 对完整会话执行说话人聚类,保持跨片段的标签一致性。"""
|
||||
try:
|
||||
|
|
@ -385,368 +352,6 @@ class AuxiliaryRuntime:
|
|||
output = self._run_embedding_pipeline(model_pipeline, audio_path)
|
||||
return self._normalize_embedding(output)
|
||||
|
||||
@staticmethod
|
||||
def _embedding_rows(result: Any, expected_count: int) -> Any:
|
||||
"""兼容 ModelScope 批量 embedding 的 Tensor、数组和逐条结果格式。"""
|
||||
import numpy as np
|
||||
|
||||
values = AuxiliaryRuntime._extract_embedding_value(result)
|
||||
if values is None:
|
||||
raise RuntimeError("speaker pipeline returned no batch embeddings")
|
||||
detach = getattr(values, "detach", None)
|
||||
if callable(detach):
|
||||
values = detach()
|
||||
cpu = getattr(values, "cpu", None)
|
||||
if callable(cpu):
|
||||
values = cpu()
|
||||
numpy_method = getattr(values, "numpy", None)
|
||||
if callable(numpy_method):
|
||||
values = numpy_method()
|
||||
try:
|
||||
array = np.asarray(values, dtype=np.float32)
|
||||
except (TypeError, ValueError):
|
||||
array = np.empty((0, 0), dtype=np.float32)
|
||||
if array.ndim == 3 and array.shape[0] == expected_count and array.shape[1] == 1:
|
||||
array = array[:, 0, :]
|
||||
if array.ndim == 1 and expected_count == 1:
|
||||
array = array.reshape(1, -1)
|
||||
if array.ndim == 2 and array.shape[0] == expected_count:
|
||||
rows = [AuxiliaryRuntime._normalize_embedding(row) for row in array]
|
||||
return rows
|
||||
|
||||
# 一些 ModelScope 版本返回每条音频一个对象,而不是一张二维矩阵。
|
||||
if isinstance(values, (list, tuple)) and len(values) == expected_count:
|
||||
rows = []
|
||||
for item in values:
|
||||
value = AuxiliaryRuntime._extract_embedding_value(item)
|
||||
if value is None:
|
||||
raise RuntimeError("speaker pipeline returned an item without an embedding")
|
||||
rows.append(AuxiliaryRuntime._normalize_embedding(value))
|
||||
return rows
|
||||
raise RuntimeError(
|
||||
"speaker pipeline batch size mismatch: "
|
||||
f"expected {expected_count}, received shape {getattr(array, 'shape', None)}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _write_pcm_window(path: str, samples: Any) -> None:
|
||||
"""把补齐后的 16kHz mono 浮点窗写成 CAM++ pipeline 可读的 PCM WAV。"""
|
||||
import numpy as np
|
||||
import wave
|
||||
|
||||
pcm = (np.clip(samples, -1.0, 1.0) * 32767.0).astype("<i2")
|
||||
with wave.open(path, "wb") as wav_file:
|
||||
wav_file.setnchannels(1)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(16000)
|
||||
wav_file.writeframes(pcm.tobytes())
|
||||
|
||||
def _extract_window_embeddings_sync(self, audio_path: str) -> list[tuple[float, float, Any]]:
|
||||
"""按 FunASR CAM++ 的 1.5s/0.75s 重叠窗批量提取声纹。"""
|
||||
import librosa
|
||||
import numpy as np
|
||||
|
||||
audio, _ = librosa.load(audio_path, sr=16000, mono=True)
|
||||
audio = np.asarray(audio, dtype=np.float32).reshape(-1)
|
||||
min_samples = int(16000 * MIN_ONLINE_SPEAKER_AUDIO_MS / 1000)
|
||||
if audio.size < min_samples:
|
||||
return []
|
||||
|
||||
window_ms = max(1500, int(os.getenv("FUNASR_SPEAKER_WINDOW_MS", "1500")))
|
||||
hop_ms = max(250, int(os.getenv("FUNASR_SPEAKER_HOP_MS", "750")))
|
||||
window_samples = int(16000 * window_ms / 1000)
|
||||
hop_samples = int(16000 * hop_ms / 1000)
|
||||
model_id = self._speaker_embedding_model_id()
|
||||
if model_id is None:
|
||||
raise RuntimeError("no loaded CAM++ speaker verification model is available")
|
||||
model_pipeline = self.models[model_id]
|
||||
|
||||
spans: list[tuple[float, float]] = []
|
||||
paths: list[str] = []
|
||||
with tempfile.TemporaryDirectory(prefix="funasr-campp-windows-") as temp_dir:
|
||||
last_end = 0
|
||||
for suggested_start in range(0, audio.size, hop_samples):
|
||||
end = min(suggested_start + window_samples, audio.size)
|
||||
if end <= last_end:
|
||||
break
|
||||
last_end = end
|
||||
# 和 FunASR sv_chunk 一样把尾窗右对齐,确保不足 1.5 秒的末尾窗
|
||||
# 尽量包含完整的新语音,而不是只用较短尾音再补大量零。
|
||||
start = max(0, end - window_samples)
|
||||
chunk = audio[start:end]
|
||||
if chunk.size < min_samples:
|
||||
break
|
||||
# 忽略近静音窗,避免把背景底噪添加成一个新说话人。
|
||||
rms = float(np.sqrt(np.mean(np.square(chunk)))) if chunk.size else 0.0
|
||||
if rms >= 0.0015:
|
||||
padded = np.zeros(window_samples, dtype=np.float32)
|
||||
padded[: chunk.size] = chunk
|
||||
path = str(Path(temp_dir) / f"window-{len(paths):04d}.wav")
|
||||
self._write_pcm_window(path, padded)
|
||||
paths.append(path)
|
||||
spans.append((start * 1000 / 16000, end * 1000 / 16000))
|
||||
if end >= audio.size:
|
||||
break
|
||||
|
||||
if not paths:
|
||||
return []
|
||||
embeddings = []
|
||||
batch_size = max(1, min(64, int(os.getenv("FUNASR_SPEAKER_BATCH_SIZE", "16"))))
|
||||
for offset in range(0, len(paths), batch_size):
|
||||
path_batch = paths[offset : offset + batch_size]
|
||||
try:
|
||||
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)
|
||||
)
|
||||
for path in path_batch
|
||||
]
|
||||
embeddings.extend(batch_embeddings)
|
||||
|
||||
return [
|
||||
(start_ms, end_ms, embedding)
|
||||
for (start_ms, end_ms), embedding in zip(spans, embeddings)
|
||||
]
|
||||
|
||||
def _map_cluster_centers(
|
||||
self,
|
||||
session_id: str,
|
||||
cluster_centers: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Map FunASR's temporary cluster labels onto stable, capped session IDs."""
|
||||
import numpy as np
|
||||
|
||||
clusters = self.speaker_clusters.setdefault(session_id, [])
|
||||
used_ids: set[int] = set()
|
||||
mapped: list[dict[str, Any]] = []
|
||||
|
||||
for raw_center in cluster_centers:
|
||||
center = self._normalize_embedding(raw_center)
|
||||
available = [
|
||||
cluster for cluster in clusters
|
||||
if int(cluster["speaker_id"]) not in used_ids
|
||||
]
|
||||
best_cluster = max(
|
||||
available,
|
||||
key=lambda cluster: float(np.dot(center, cluster["embedding"])),
|
||||
default=None,
|
||||
)
|
||||
best_score = (
|
||||
float(np.dot(center, best_cluster["embedding"]))
|
||||
if best_cluster is not None
|
||||
else -1.0
|
||||
)
|
||||
matched = (
|
||||
best_cluster is not None
|
||||
and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD
|
||||
)
|
||||
created = False
|
||||
|
||||
if not matched and len(clusters) < ONLINE_MAX_SPEAKERS:
|
||||
speaker_id = len(clusters)
|
||||
clusters.append(
|
||||
{"speaker_id": speaker_id, "embedding": center, "count": 1}
|
||||
)
|
||||
confidence = 0.75
|
||||
strategy = "online_embedding_cluster_new"
|
||||
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.
|
||||
best_cluster = max(
|
||||
clusters,
|
||||
key=lambda cluster: float(np.dot(center, cluster["embedding"])),
|
||||
)
|
||||
best_score = float(np.dot(center, best_cluster["embedding"]))
|
||||
speaker_id = int(best_cluster["speaker_id"])
|
||||
confidence = max(0.0, best_score)
|
||||
if matched:
|
||||
count = int(best_cluster["count"])
|
||||
weight = 1.0 / min(count + 1, 20)
|
||||
best_cluster["embedding"] = self._normalize_embedding(
|
||||
best_cluster["embedding"] * (1.0 - weight)
|
||||
+ center * weight
|
||||
)
|
||||
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)
|
||||
mapped.append(
|
||||
{
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_confidence": round(
|
||||
max(0.6, min(1.0, confidence)), 3
|
||||
),
|
||||
"speaker_strategy": strategy,
|
||||
}
|
||||
)
|
||||
return mapped
|
||||
|
||||
def _cluster_speaker_history(
|
||||
self,
|
||||
session_id: str,
|
||||
) -> tuple[list[list[float]], list[dict[str, Any]]]:
|
||||
"""Re-cluster the rolling CAM++ history with FunASR's own backend."""
|
||||
import numpy as np
|
||||
import torch
|
||||
from funasr.models.campplus.cluster_backend import ClusterBackend
|
||||
from funasr.models.campplus.utils import postprocess
|
||||
|
||||
history = self.speaker_history[session_id]
|
||||
embeddings = torch.as_tensor(
|
||||
np.stack(list(history["embeddings"])),
|
||||
dtype=torch.float32,
|
||||
device="cpu",
|
||||
)
|
||||
if self._speaker_cluster_backend is None:
|
||||
self._speaker_cluster_backend = ClusterBackend(
|
||||
merge_thr=ONLINE_CLUSTER_MERGE_THRESHOLD
|
||||
).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.
|
||||
labels = self._speaker_cluster_backend(embeddings, oracle_num=None)
|
||||
labels = np.asarray(labels)
|
||||
chunks = [
|
||||
[start_ms / 1000.0, end_ms / 1000.0, None]
|
||||
for start_ms, end_ms in history["chunks"]
|
||||
]
|
||||
segments, centers = self._speaker_postprocess(
|
||||
chunks,
|
||||
None,
|
||||
labels,
|
||||
embeddings,
|
||||
return_spk_center=True,
|
||||
)
|
||||
stable_clusters = self._map_cluster_centers(session_id, centers)
|
||||
return segments, stable_clusters
|
||||
|
||||
def _assign_embedding(
|
||||
self,
|
||||
session_id: str,
|
||||
embedding: Any,
|
||||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> dict[str, Any]:
|
||||
"""Assign a fallback whole-turn embedding with FunASR's 15-speaker policy."""
|
||||
import numpy as np
|
||||
|
||||
embedding = self._normalize_embedding(embedding)
|
||||
clusters = self.speaker_clusters.setdefault(session_id, [])
|
||||
best_cluster: dict[str, Any] | None = None
|
||||
best_score = -1.0
|
||||
for cluster in clusters:
|
||||
if embedding.shape != cluster["embedding"].shape:
|
||||
raise RuntimeError("speaker embedding dimension changed within the session")
|
||||
score = float(np.dot(embedding, cluster["embedding"]))
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_cluster = cluster
|
||||
|
||||
if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD:
|
||||
count = int(best_cluster["count"])
|
||||
weight = 1.0 / min(count + 1, 20)
|
||||
best_cluster["embedding"] = self._normalize_embedding(
|
||||
best_cluster["embedding"] * (1.0 - weight)
|
||||
+ embedding * weight
|
||||
)
|
||||
best_cluster["count"] = count + 1
|
||||
speaker_id = int(best_cluster["speaker_id"])
|
||||
confidence = best_score
|
||||
strategy = "online_embedding_cluster_match"
|
||||
elif len(clusters) < ONLINE_MAX_SPEAKERS:
|
||||
speaker_id = len(clusters)
|
||||
clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1})
|
||||
confidence = 0.75
|
||||
strategy = "online_embedding_cluster_new"
|
||||
else:
|
||||
# Match FunASR's fallback after the identity limit is reached.
|
||||
speaker_id = int(best_cluster["speaker_id"]) if best_cluster else 0
|
||||
confidence = max(0.0, best_score)
|
||||
strategy = "online_embedding_cluster_limit_fallback"
|
||||
|
||||
return {
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_name": f"说话人 {speaker_id + 1}",
|
||||
"speaker_evidence": "fresh",
|
||||
"speaker_confidence": round(max(0.6, min(1.0, confidence)), 3),
|
||||
"speaker_strategy": strategy,
|
||||
"speaker_status": "confirmed",
|
||||
"speaker_reason": "CAM++ rolling history was clustered with FunASR's speaker backend",
|
||||
"start_time": start_time_ms,
|
||||
"end_time": end_time_ms,
|
||||
}
|
||||
|
||||
async def track_speakers(
|
||||
self,
|
||||
audio_path: str,
|
||||
session_id: str,
|
||||
start_time_ms: float,
|
||||
end_time_ms: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Track a completed VAD turn using FunASR's rolling-window clusterer."""
|
||||
async with self.inference_lock:
|
||||
now = time.monotonic()
|
||||
for stale_id, seen in list(self.speaker_last_seen.items()):
|
||||
if now - seen > 1800:
|
||||
self.reset_speaker_session(stale_id)
|
||||
self.speaker_last_seen[session_id] = now
|
||||
windows = await asyncio.to_thread(self._extract_window_embeddings_sync, audio_path)
|
||||
if session_id not in self.speaker_last_seen or not windows:
|
||||
return []
|
||||
|
||||
history = self.speaker_history.get(session_id)
|
||||
if history is None:
|
||||
history = {
|
||||
"chunks": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS),
|
||||
"embeddings": deque(maxlen=ONLINE_SPEAKER_HISTORY_CHUNKS),
|
||||
}
|
||||
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)
|
||||
)
|
||||
history["embeddings"].append(embedding.copy())
|
||||
|
||||
clustered_segments, stable_clusters = self._cluster_speaker_history(session_id)
|
||||
results: list[dict[str, Any]] = []
|
||||
for segment_start, segment_end, cluster_id in clustered_segments:
|
||||
segment_start_ms = max(start_time_ms, float(segment_start) * 1000.0)
|
||||
segment_end_ms = min(end_time_ms, float(segment_end) * 1000.0)
|
||||
if segment_end_ms <= segment_start_ms:
|
||||
continue
|
||||
cluster_index = int(cluster_id)
|
||||
if cluster_index < 0 or cluster_index >= len(stable_clusters):
|
||||
continue
|
||||
stable = stable_clusters[cluster_index]
|
||||
speaker_id = int(stable["speaker_id"])
|
||||
results.append(
|
||||
{
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_name": f"说话人 {speaker_id + 1}",
|
||||
"speaker_evidence": "fresh",
|
||||
"speaker_confidence": stable["speaker_confidence"],
|
||||
"speaker_strategy": stable["speaker_strategy"],
|
||||
"speaker_status": "confirmed",
|
||||
"speaker_reason": "CAM++ rolling history was clustered with FunASR's speaker backend",
|
||||
"start_time": segment_start_ms,
|
||||
"end_time": segment_end_ms,
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
async def resolve_speaker(
|
||||
self,
|
||||
audio_path: str,
|
||||
|
|
@ -766,17 +371,53 @@ class AuxiliaryRuntime:
|
|||
if embedding is None:
|
||||
return {"speaker_id": -1, "speaker_evidence": "pending", "speaker_confidence": 0.0,
|
||||
"speaker_status": "insufficient_audio", "speaker_reason": "音频不足 800ms,未提取声纹"}
|
||||
import numpy as np
|
||||
|
||||
# reset 可能在模型线程运行期间到达;结束的会话不能被晚到结果重新创建。
|
||||
if session_id not in self.speaker_last_seen:
|
||||
return None
|
||||
return self._assign_embedding(
|
||||
session_id, embedding, start_time_ms, end_time_ms
|
||||
embedding = self._normalize_embedding(embedding)
|
||||
clusters = self.speaker_clusters.setdefault(session_id, [])
|
||||
best_cluster: dict[str, Any] | None = None
|
||||
best_score = -1.0
|
||||
for cluster in clusters:
|
||||
if embedding.shape != cluster["embedding"].shape:
|
||||
raise RuntimeError("speaker embedding dimension changed within the session")
|
||||
score = float(np.dot(embedding, cluster["embedding"]))
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_cluster = cluster
|
||||
|
||||
if best_cluster is not None and best_score >= ONLINE_SPEAKER_MATCH_THRESHOLD:
|
||||
count = int(best_cluster["count"])
|
||||
best_cluster["embedding"] = self._normalize_embedding(
|
||||
(best_cluster["embedding"] * count) + embedding
|
||||
)
|
||||
best_cluster["count"] = count + 1
|
||||
speaker_id = int(best_cluster["speaker_id"])
|
||||
confidence = best_score
|
||||
strategy = "online_embedding_cluster_match"
|
||||
else:
|
||||
speaker_id = len(clusters)
|
||||
clusters.append({"speaker_id": speaker_id, "embedding": embedding, "count": 1})
|
||||
confidence = 0.75
|
||||
strategy = "online_embedding_cluster_new"
|
||||
|
||||
return {
|
||||
"speaker_id": speaker_id,
|
||||
"speaker_name": f"说话人 {speaker_id + 1}",
|
||||
"speaker_evidence": "fresh",
|
||||
"speaker_confidence": round(max(0.6, min(1.0, confidence)), 3),
|
||||
"speaker_strategy": strategy,
|
||||
"speaker_status": "confirmed",
|
||||
"speaker_reason": "当前片段独立声纹已完成在线聚类",
|
||||
"start_time": start_time_ms,
|
||||
"end_time": end_time_ms,
|
||||
}
|
||||
|
||||
def reset_speaker_session(self, session_id: str) -> None:
|
||||
"""释放已结束 WebSocket 的聚类中心,防止长时间运行时内存增长。"""
|
||||
self.speaker_clusters.pop(session_id, None)
|
||||
self.speaker_history.pop(session_id, None)
|
||||
self.speaker_last_seen.pop(session_id, None)
|
||||
|
||||
|
||||
|
|
@ -792,8 +433,9 @@ 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.
|
||||
ready = speaker_model_id is not None
|
||||
# ready 表示实时链路的两个核心模型都可用;完整 diarization 是否
|
||||
# 预加载不影响这里的结果。
|
||||
ready = vad_model_id is not None and speaker_model_id is not None
|
||||
return web.json_response(
|
||||
{
|
||||
"ready": ready,
|
||||
|
|
@ -803,10 +445,6 @@ async def health_handler(request: web.Request) -> web.Response:
|
|||
"device": AUXILIARY_DEVICE,
|
||||
"speaker_embedding_model": speaker_model_id,
|
||||
"speaker_embedding_ready": speaker_model_id is not None,
|
||||
"punctuation_ready": any(
|
||||
config.get("kind") == "punctuation" and model_id in runtime.models
|
||||
for model_id, config in runtime.assets.items()
|
||||
),
|
||||
"models": runtime.status,
|
||||
}
|
||||
)
|
||||
|
|
@ -831,27 +469,6 @@ async def vad_handler(request: web.Request) -> web.Response:
|
|||
Path(temp_path).unlink(missing_ok=True)
|
||||
|
||||
|
||||
async def punctuation_handler(request: web.Request) -> web.Response:
|
||||
"""Return punctuation when the optional local CT-Transformer is available."""
|
||||
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
|
||||
try:
|
||||
payload = await request.json()
|
||||
except (ValueError, web.HTTPException):
|
||||
return web.json_response({"error": "request body must be JSON"}, status=400)
|
||||
text = payload.get("text") if isinstance(payload, dict) else None
|
||||
if not isinstance(text, str):
|
||||
return web.json_response({"error": "JSON field 'text' must be a string"}, status=400)
|
||||
if not text.strip():
|
||||
return web.json_response({"text": text, "available": True})
|
||||
try:
|
||||
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.
|
||||
print(f"[punctuation] unavailable: {exc}", flush=True)
|
||||
return web.json_response({"text": text, "available": False, "error": str(exc)})
|
||||
|
||||
|
||||
def _normalize_diarization_segments(result: Any) -> list[dict[str, Any]]:
|
||||
"""将不同 ModelScope 版本的聚类输出统一为 start/end/speaker 字段。"""
|
||||
# ModelScope 通常返回 {'text': [[start_sec, end_sec, speaker_id], ...]};
|
||||
|
|
@ -962,43 +579,6 @@ async def speaker_resolve_handler(request: web.Request) -> web.Response:
|
|||
Path(temp_path).unlink(missing_ok=True)
|
||||
|
||||
|
||||
async def speaker_track_handler(request: web.Request) -> web.Response:
|
||||
"""接收一个 FunASR turn,返回经过短段平滑的 CAM++ 滑窗标签。"""
|
||||
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
|
||||
form = await request.post()
|
||||
upload = form.get("file")
|
||||
if not isinstance(upload, FileField):
|
||||
return web.json_response({"error": "multipart field 'file' is required"}, status=400)
|
||||
session_id = str(form.get("session_id") or "").strip()
|
||||
if not session_id:
|
||||
return web.json_response({"error": "multipart field 'session_id' is required"}, status=400)
|
||||
try:
|
||||
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)
|
||||
except ValueError:
|
||||
return web.json_response({"error": "turn time fields must be numbers"}, status=400)
|
||||
|
||||
temp_path: str | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as temp_file:
|
||||
temp_file.write(upload.file.read())
|
||||
temp_path = temp_file.name
|
||||
segments = await runtime.track_speakers(
|
||||
temp_path, session_id, start_time_ms, end_time_ms
|
||||
)
|
||||
return web.json_response({"segments": segments})
|
||||
except Exception as exc:
|
||||
print(
|
||||
f"[speaker] track failed: session_id={session_id}, "
|
||||
f"start={start_time_ms}, end={end_time_ms}, error={exc}",
|
||||
flush=True,
|
||||
)
|
||||
return web.json_response({"error": str(exc)}, status=500)
|
||||
finally:
|
||||
if temp_path:
|
||||
Path(temp_path).unlink(missing_ok=True)
|
||||
|
||||
|
||||
async def speaker_reset_handler(request: web.Request) -> web.Response:
|
||||
"""释放已经结束的实时会话聚类中心。"""
|
||||
runtime: AuxiliaryRuntime = request.app[MODEL_SERVICE_KEY]
|
||||
|
|
@ -1017,10 +597,8 @@ async def create_app() -> web.Application:
|
|||
app[MODEL_SERVICE_KEY] = runtime
|
||||
app.router.add_get("/health", health_handler)
|
||||
app.router.add_post("/v1/vad", vad_handler)
|
||||
app.router.add_post("/v1/punctuation", punctuation_handler)
|
||||
app.router.add_post("/v1/diarization", diarization_handler)
|
||||
app.router.add_post("/v1/speaker/resolve", speaker_resolve_handler)
|
||||
app.router.add_post("/v1/speaker/track", speaker_track_handler)
|
||||
app.router.add_post("/v1/speaker/reset", speaker_reset_handler)
|
||||
return app
|
||||
|
||||
|
|
@ -1,12 +1,229 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Download models using the project's shared model manifest."""
|
||||
"""为独立服务部署下载 ASR 和配套辅助模型。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# 同时支持直接执行脚本和 `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:
|
||||
"""检查 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-download.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, diarization, and aligner assets",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
manifest = load_manifest()
|
||||
models_dir = args.models_dir.resolve()
|
||||
cache_dir = args.cache_dir.resolve() if args.cache_dir else None
|
||||
selected_assets: list[tuple[str, dict[str, object]]] = []
|
||||
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,29 +1,54 @@
|
|||
"""Compatibility exports for shared model download scripts."""
|
||||
"""独立 VLLM 部署项目的模型清单与路径解析辅助函数。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
|
||||
# Direct script execution puts scripts/ first on sys.path; add the repository root
|
||||
# so the shared backend manifest remains importable in both launch modes.
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
MANIFEST_PATH = PROJECT_ROOT / "model_manifest.json"
|
||||
|
||||
from backend.model_manifest import (
|
||||
MANIFEST_PATH,
|
||||
auxiliary_models,
|
||||
load_manifest,
|
||||
model_directory,
|
||||
resolve_auxiliary_model_id,
|
||||
resolve_model_id,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"MANIFEST_PATH",
|
||||
"PROJECT_ROOT",
|
||||
"auxiliary_models",
|
||||
"load_manifest",
|
||||
"model_directory",
|
||||
"resolve_auxiliary_model_id",
|
||||
"resolve_model_id",
|
||||
]
|
||||
def load_manifest(path: Path = MANIFEST_PATH) -> dict[str, Any]:
|
||||
"""读取本项目自己的模型清单,整个过程不导入原项目代码。"""
|
||||
with path.open("r", encoding="utf-8") as manifest_file:
|
||||
manifest = json.load(manifest_file)
|
||||
if not isinstance(manifest.get("models"), dict) or not manifest["models"]:
|
||||
raise ValueError("model_manifest.json 必须包含非空的 models 对象")
|
||||
return manifest
|
||||
|
||||
|
||||
def resolve_model_id(model: str | None, manifest: dict[str, Any]) -> str:
|
||||
"""将默认值、短别名或完整模型 ID 解析为一个 ASR 模型。"""
|
||||
models = manifest["models"]
|
||||
requested = (model or "default").strip()
|
||||
if requested == "default":
|
||||
requested = str(manifest["default_model"])
|
||||
|
||||
if requested in models:
|
||||
return requested
|
||||
|
||||
for model_id, config in models.items():
|
||||
if requested.lower() == str(config.get("alias", "")).lower():
|
||||
return model_id
|
||||
raise ValueError(f"不支持的 ASR 模型 '{model}',可选模型:{', '.join(models)}")
|
||||
|
||||
|
||||
def model_directory(model_id: str, manifest: dict[str, Any], models_dir: Path) -> Path:
|
||||
"""根据清单返回 ASR 或辅助模型实际使用的本地目录。"""
|
||||
config = manifest.get("models", {}).get(model_id)
|
||||
if config is None:
|
||||
config = manifest.get("auxiliary_models", {}).get(model_id)
|
||||
if not isinstance(config, dict) or not config.get("directory"):
|
||||
raise ValueError(f"模型 '{model_id}' 在清单中没有配置本地目录")
|
||||
return models_dir / str(config["directory"])
|
||||
|
||||
|
||||
def auxiliary_models(manifest: dict[str, Any]) -> dict[str, dict[str, Any]]:
|
||||
"""返回可独立部署的 VAD、说话人和对齐模型资产。"""
|
||||
models = manifest.get("auxiliary_models", {})
|
||||
if not isinstance(models, dict):
|
||||
raise ValueError("model_manifest.json 的 auxiliary_models 必须是对象")
|
||||
return {str(model_id): config for model_id, config in models.items() if isinstance(config, dict)}
|
||||
|
|
|
|||
|
|
@ -19,14 +19,16 @@ 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
|
||||
try:
|
||||
from .model_manifest import load_manifest, model_directory, resolve_model_id
|
||||
except ImportError:
|
||||
from model_manifest import load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
def has_model_weights(model_path: Path) -> bool:
|
||||
|
|
@ -123,7 +125,7 @@ def build_server_command(args: argparse.Namespace, model_id: str, model_path: Pa
|
|||
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")
|
||||
raise RuntimeError(f"{executable_name} was not found; install requirements-deploy.txt first")
|
||||
|
||||
served_model_name = args.served_model_name or model_id
|
||||
command = [
|
||||
|
|
@ -7,7 +7,7 @@ import numpy as np
|
|||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from backend.auxiliary_server import (
|
||||
from scripts.auxiliary_server import (
|
||||
AuxiliaryRuntime,
|
||||
_coerce_finite_float,
|
||||
_normalize_diarization_segments,
|
||||
|
|
@ -42,8 +42,8 @@ class AuxiliaryServerTests(unittest.TestCase):
|
|||
"aligner": {"kind": "forced_aligner"},
|
||||
}
|
||||
loaded_kinds = []
|
||||
with patch("backend.auxiliary_server._asset_ready", return_value=True), \
|
||||
patch("backend.auxiliary_server.model_directory", return_value=runtime.manifest and SimpleNamespace()), \
|
||||
with patch("scripts.auxiliary_server._asset_ready", return_value=True), \
|
||||
patch("scripts.auxiliary_server.model_directory", return_value=runtime.manifest and SimpleNamespace()), \
|
||||
patch.object(runtime, "_load_asset", side_effect=lambda model_id, config, path: loaded_kinds.append(config["kind"]) or object()):
|
||||
runtime.preload()
|
||||
self.assertEqual(loaded_kinds, ["vad"])
|
||||
|
|
@ -57,8 +57,8 @@ class AuxiliaryServerTests(unittest.TestCase):
|
|||
"vad": {"kind": "vad"},
|
||||
"campplus": {"kind": "speaker_verification"},
|
||||
}
|
||||
with patch("backend.auxiliary_server._asset_ready", side_effect=[True, False]), \
|
||||
patch("backend.auxiliary_server.model_directory", return_value=SimpleNamespace()):
|
||||
with patch("scripts.auxiliary_server._asset_ready", side_effect=[True, False]), \
|
||||
patch("scripts.auxiliary_server.model_directory", return_value=SimpleNamespace()):
|
||||
with self.assertRaisesRegex(RuntimeError, "campplus"):
|
||||
runtime.preload()
|
||||
|
||||
|
|
@ -76,8 +76,8 @@ class AuxiliaryServerTests(unittest.TestCase):
|
|||
"""VAD 是核心依赖,缺失时错误必须给出可执行的修复方向。"""
|
||||
runtime = AuxiliaryRuntime()
|
||||
runtime.assets = {"vad": {"kind": "vad"}}
|
||||
with patch("backend.auxiliary_server._asset_ready", return_value=False), \
|
||||
patch("backend.auxiliary_server.model_directory", return_value=SimpleNamespace(__str__=lambda self: "/models/vad")):
|
||||
with patch("scripts.auxiliary_server._asset_ready", return_value=False), \
|
||||
patch("scripts.auxiliary_server.model_directory", return_value=SimpleNamespace(__str__=lambda self: "/models/vad")):
|
||||
with self.assertRaisesRegex(RuntimeError, "download_models.py --auxiliary-only"):
|
||||
runtime.preload()
|
||||
|
||||
|
|
@ -8,7 +8,7 @@ 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
|
||||
from scripts.model_manifest import auxiliary_models, load_manifest, model_directory, resolve_model_id
|
||||
|
||||
|
||||
class ModelManifestTests(unittest.TestCase):
|
||||
|
|
@ -22,13 +22,10 @@ class ModelManifestTests(unittest.TestCase):
|
|||
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"])
|
||||
))
|
||||
def test_manifest_has_only_asr_models(self) -> None:
|
||||
self.assertEqual(
|
||||
resolve_model_id("paraformer-zh-streaming", self.manifest),
|
||||
"iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online",
|
||||
set(self.manifest["models"]),
|
||||
{"Qwen/Qwen3-ASR-0.6B", "Qwen/Qwen3-ASR-1.7B"},
|
||||
)
|
||||
|
||||
def test_manifest_has_auxiliary_runtime_assets(self) -> None:
|
||||
|
|
@ -39,7 +36,7 @@ class ModelManifestTests(unittest.TestCase):
|
|||
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"
|
||||
models_dir = Path(__file__).resolve().parents[1] / "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))
|
||||
|
||||
|
|
@ -7,7 +7,7 @@ import unittest
|
|||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from backend.serve_qwen_legacy import SERVER_PORT, build_parser, build_server_command
|
||||
from scripts.serve import SERVER_PORT, build_parser, build_server_command
|
||||
|
||||
|
||||
class ServeConfigTests(unittest.TestCase):
|
||||
|
|
@ -29,7 +29,7 @@ class ServeConfigTests(unittest.TestCase):
|
|||
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")
|
||||
@patch("scripts.serve.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([])
|
||||
Loading…
Reference in New Issue