Use FunASR native realtime WebSocket flow
parent
fbaf56fd9f
commit
13278d3aea
|
|
@ -1,30 +1,32 @@
|
|||
# Local model assets; the backend never downloads models at startup.
|
||||
# Local model assets; model loading never downloads weights at startup.
|
||||
MODEL_DIR=models
|
||||
# Set this if CAM++ is not at a model_manifest.json path.
|
||||
# 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
|
||||
|
||||
# FunASR streaming ASR model. The backend launcher requires it to exist locally.
|
||||
# Local FunASR streaming ASR and VAD models.
|
||||
FUNASR_ASR_MODEL=paraformer-zh-streaming
|
||||
FUNASR_VAD_MODEL=fsmn-vad
|
||||
|
||||
# ASR and VAD may use different devices to limit peak GPU memory.
|
||||
FUNASR_DEVICE=cuda:0
|
||||
FUNASR_VAD_DEVICE=cpu
|
||||
|
||||
# [left context, current chunk, right lookahead]; 10 * 60ms = 600ms.
|
||||
# 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_VAD_CHUNK_MS=200
|
||||
FUNASR_MAX_SEGMENT_SEC=30
|
||||
FUNASR_FINALIZE_TIMEOUT_SECONDS=300
|
||||
FUNASR_NATIVE_WS_HOST=127.0.0.1
|
||||
FUNASR_NATIVE_WS_PORT=10095
|
||||
|
||||
# CAM++ is required. The backend launcher starts this service itself.
|
||||
# CAM++ is required and started by scripts/run_backend.py.
|
||||
AUXILIARY_SERVICE_URL=http://127.0.0.1:8010
|
||||
AUXILIARY_DEVICE=cpu
|
||||
AUXILIARY_PRELOAD_KINDS=speaker_verification
|
||||
|
||||
# Standalone frontend and browser-facing backend origin.
|
||||
# Frontend and backend run independently. The frontend proxies /ws and /api/stop.
|
||||
FRONTEND_HOST=127.0.0.1
|
||||
FRONTEND_PORT=8080
|
||||
FRONTEND_ORIGIN=http://127.0.0.1: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
|
||||
|
|
|
|||
|
|
@ -1,59 +1,43 @@
|
|||
# FunASR realtime browser demo
|
||||
|
||||
The frontend and backend start separately. The backend launcher loads three
|
||||
required local assets: streaming ASR, FSMN VAD, and CAM++ speaker verification.
|
||||
It starts the CAM++ model service and the WebSocket service and stops both
|
||||
together. No model is downloaded by the launcher. Use the shared model manifest
|
||||
and downloader to prepare the three required snapshots:
|
||||
The frontend and backend start separately. The backend starts these supervised processes:
|
||||
|
||||
~~~powershell
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
- 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 the assets under models/ or set MODEL_DIR in .env. The default
|
||||
FunASR names resolve to these local 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 damo/ variant)
|
||||
- `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 your directories have different names, set FUNASR_ASR_MODEL and
|
||||
FUNASR_VAD_MODEL to their paths. The CAM++ directories follow
|
||||
model_manifest.json; set CAM_MODEL_PATH for another location. Backend startup reports every checked path when a
|
||||
required asset is missing.
|
||||
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.
|
||||
|
||||
## Start
|
||||
## Install and start
|
||||
|
||||
First install a torch/torchaudio build suitable for the host CPU or CUDA,
|
||||
then install the project dependencies:
|
||||
Install a torch/torchaudio build suitable for the host, then install the project dependencies and models:
|
||||
|
||||
~~~powershell
|
||||
cd D:\github-project\ASR\Asr-demo
|
||||
python -m pip install -r requirements-funasr.txt
|
||||
python -m pip install -r requirements-auxiliary.txt
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
if (-not (Test-Path .env)) { Copy-Item .env.funasr.example .env }
|
||||
~~~
|
||||
|
||||
In one terminal start the backend:
|
||||
Start the backend and frontend in separate terminals:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_backend.py
|
||||
~~~
|
||||
|
||||
In another terminal start the frontend:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_frontend.py
|
||||
~~~
|
||||
|
||||
Open http://127.0.0.1:8080/. The backend WebSocket listens on port 8082 and
|
||||
the CAM++ service on port 8010. BACKEND_PUBLIC_URL must be reachable from
|
||||
the browser. Set FRONTEND_ORIGIN to the exact frontend origin if it differs
|
||||
from the default.
|
||||
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}`.
|
||||
|
||||
The browser sends start, PCM16 frames, and stop/eof over WebSocket. Speaker
|
||||
labels are required. Short or unusable speech may still receive an unknown
|
||||
speaker label, but missing CAM++ prevents backend startup.
|
||||
`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++.
|
||||
|
|
|
|||
13
README.md
13
README.md
|
|
@ -1,20 +1,16 @@
|
|||
# FunASR realtime ASR demo
|
||||
|
||||
This branch runs streaming ASR, VAD, and CAM++ from local model assets.
|
||||
The browser UI starts as a separate service.
|
||||
This branch runs FunASR's native online WebSocket flow with local streaming ASR and FSMN VAD models. CAM++ speaker embeddings are handled by the separate auxiliary service. The browser page, JavaScript, and CSS are copied byte-for-byte from the local `tencent-demo/static` directory.
|
||||
|
||||
## Start
|
||||
|
||||
Install a torch/torchaudio build for the host, then install
|
||||
requirements-funasr.txt and requirements-auxiliary.txt. Copy
|
||||
.env.funasr.example to .env if no .env exists, then download the three
|
||||
required FunASR models with the shared project manifest:
|
||||
Install a torch/torchaudio build for the host, then install `requirements-funasr.txt` and `requirements-auxiliary.txt`. Copy `.env.funasr.example` to `.env` if no `.env` exists, then download the three required FunASR assets:
|
||||
|
||||
~~~powershell
|
||||
python scripts/download_models.py --funasr-runtime
|
||||
~~~
|
||||
|
||||
Start the backend (CAM++ model service plus WebSocket) in one terminal:
|
||||
Start the backend (CAM++, native FunASR WSS, and the browser protocol bridge) in one terminal:
|
||||
|
||||
~~~powershell
|
||||
python scripts\run_backend.py
|
||||
|
|
@ -26,5 +22,4 @@ Start the frontend in another terminal:
|
|||
python scripts\run_frontend.py
|
||||
~~~
|
||||
|
||||
Open http://127.0.0.1:8080/. Model directory layout and address settings
|
||||
are described in FUNASR_README.md.
|
||||
Open the URL configured by `FRONTEND_HOST` and `FRONTEND_PORT`. The frontend keeps the Tencent demo's same-origin `/ws` and `/api/stop` calls and proxies them to `BACKEND_INTERNAL_URL`. The public backend bridge defaults to port 8082; the native FunASR WS socket defaults to loopback port 10095. Model paths and device settings are described in [FUNASR_README.md](FUNASR_README.md).
|
||||
|
|
|
|||
|
|
@ -0,0 +1,888 @@
|
|||
"""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 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}
|
||||
websocket.status_dict_vad = {"cache": {}, "is_final": False}
|
||||
websocket.status_dict_punc = {"cache": {}}
|
||||
|
||||
websocket.chunk_interval = 10
|
||||
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 "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
|
||||
|
||||
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
|
||||
|
||||
if rec_result.get("text"):
|
||||
mode = "2pass-online" if "2pass" in (websocket.mode or "") else websocket.mode
|
||||
message = {
|
||||
"mode": mode,
|
||||
"text": rec_result["text"],
|
||||
"wav_name": websocket.wav_name,
|
||||
"is_final": bool(
|
||||
websocket.status_dict_asr_online.get("is_final", False) or (not websocket.is_speaking)
|
||||
),
|
||||
}
|
||||
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
|
||||
File diff suppressed because it is too large
Load Diff
|
|
@ -4,3 +4,4 @@ funasr==1.4.16
|
|||
modelscope[framework]==1.34.0
|
||||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
websockets>=12,<14
|
||||
|
|
|
|||
|
|
@ -1,8 +1,5 @@
|
|||
// ===== 页面元素 =====
|
||||
// ===== 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');
|
||||
|
|
@ -22,13 +19,13 @@ const elMicStatus = document.getElementById('micStatus');
|
|||
const elMicTimer = document.getElementById('micTimer');
|
||||
const elMicElapsed = document.getElementById('micElapsed');
|
||||
|
||||
// 输入模式标签页
|
||||
// Input mode tabs
|
||||
const elTabMic = document.getElementById('tabMic');
|
||||
const elTabFile = document.getElementById('tabFile');
|
||||
const elPanelMic = document.getElementById('panelMic');
|
||||
const elPanelFile = document.getElementById('panelFile');
|
||||
|
||||
// 文件选择区域
|
||||
// File upload
|
||||
const elAudioFile = document.getElementById('audioFile');
|
||||
const elFileInfo = document.getElementById('fileInfo');
|
||||
const elAudioMeta = document.getElementById('audioMeta');
|
||||
|
|
@ -39,12 +36,12 @@ const elSpeedControl = document.getElementById('speedControl');
|
|||
const elSpeedSlider = document.getElementById('speedSlider');
|
||||
const elSpeedValue = document.getElementById('speedValue');
|
||||
|
||||
// ===== 说话人分离开关 =====
|
||||
// ===== Speaker Diarization Toggle =====
|
||||
elSpeakerDiarization.addEventListener('change', () => {
|
||||
elDiarizationLabel.textContent = elSpeakerDiarization.checked ? '开启' : '关闭';
|
||||
});
|
||||
|
||||
// ===== 日志区域 =====
|
||||
// ===== Log Area =====
|
||||
elBtnClearLog.addEventListener('click', () => { elLogArea.innerHTML = ''; });
|
||||
|
||||
function appendLog(msg) {
|
||||
|
|
@ -55,20 +52,12 @@ function appendLog(msg) {
|
|||
const typeClass = 'log-type-' + (msg.type || 'unknown');
|
||||
const entry = document.createElement('div');
|
||||
entry.className = 'log-entry';
|
||||
// 原始文本不作为 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);
|
||||
entry.innerHTML = `<span class="log-time">${ts}</span><span class="${typeClass}">${JSON.stringify(msg)}</span>`;
|
||||
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;
|
||||
|
|
@ -82,7 +71,7 @@ let micWorklet = null;
|
|||
let micTimerInterval = null;
|
||||
let micStartTime = 0;
|
||||
|
||||
// 输入模式(麦克风 / 文件)
|
||||
// Input mode (mic / file)
|
||||
let inputMode = 'mic';
|
||||
let selectedFile = null;
|
||||
|
||||
|
|
@ -91,28 +80,25 @@ const EXT_FORMAT_MAP = {
|
|||
'pcm': 1, 'wav': 12, 'mp3': 8, 'm4a': 14,
|
||||
'aac': 16, 'opus': 10, 'ogg': 10, 'silk': 6, 'speex': 4
|
||||
};
|
||||
// PCM/WAV 的默认发送倍速;实时验证默认按 1 倍速输入。
|
||||
// 不同格式的默认发送倍速:PCM/WAV 实时速度 1x,压缩格式解压快可提速
|
||||
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;
|
||||
// 实时 WebSocket 需要服务端逐帧读取音频;压缩格式必须等文件完整后才能解码,
|
||||
// 因此本次流式验证只允许 PCM/WAV,避免把整段上传伪装成实时识别。
|
||||
const STREAMABLE_AUDIO_EXTENSIONS = new Set(['pcm', 'wav']);
|
||||
const UNSUPPORTED_STREAMING_EXTENSIONS = new Set(['m4a']);
|
||||
|
||||
const SPEAKER_COLORS = ['#4a7dff', '#52c41a', '#faad14', '#ff4d4f', '#9254de', '#13c2c2'];
|
||||
|
||||
let sentenceMap = {};
|
||||
let speakerOrderMap = {};
|
||||
let speakerOrderCounter = 0;
|
||||
let displayStateSupported = false;
|
||||
let displayRevision = -1;
|
||||
let lastConfirmedBubbleEl = null;
|
||||
let pendingSpanMap = {};
|
||||
|
||||
// ===== 输入模式标签页 =====
|
||||
// ===== Input Mode Tabs =====
|
||||
function switchMode(mode) {
|
||||
if (ws) return;
|
||||
inputMode = mode;
|
||||
elTabMic.classList.toggle('active', mode === 'mic');
|
||||
elTabFile.classList.toggle('active', mode === 'file');
|
||||
|
|
@ -127,20 +113,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 (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
selectedFile = null;
|
||||
e.target.value = '';
|
||||
elFileInfo.textContent = '实时测试只支持 PCM 或 WAV,请先转换音频格式';
|
||||
elFileInfo.textContent = 'M4A 暂不支持直接上传,请先转成 WAV 或 MP3';
|
||||
elFileInfo.classList.remove('has-file');
|
||||
elAudioMeta.style.display = 'none';
|
||||
elSpeedControl.style.display = 'none';
|
||||
elBtnStart.disabled = true;
|
||||
showToast('压缩音频不能按当前实时 WebSocket 逐帧识别,请转成 PCM 或 WAV', true);
|
||||
showToast('M4A 容器格式无法按当前实时切片方式直接识别,请转成 WAV 或 MP3', true);
|
||||
return;
|
||||
}
|
||||
selectedFile = file;
|
||||
|
|
@ -151,12 +137,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() : '';
|
||||
|
|
@ -221,7 +207,7 @@ async function parseAudioMeta(file) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== 复制和提示 =====
|
||||
// ===== Copy & Toast =====
|
||||
function showToast(message, isError) {
|
||||
const toast = document.createElement('div');
|
||||
toast.className = 'copy-toast' + (isError ? ' copy-toast-error' : '');
|
||||
|
|
@ -246,7 +232,7 @@ function handleCopyClick(btn, textEl) {
|
|||
|
||||
elBtnCopyVoiceId.addEventListener('click', () => handleCopyClick(elBtnCopyVoiceId, elVoiceIdDisplay));
|
||||
|
||||
// ===== WAV 导出 =====
|
||||
// ===== WAV Export =====
|
||||
function buildWavBlob(pcmChunks) {
|
||||
let totalLen = 0;
|
||||
for (const c of pcmChunks) totalLen += c.byteLength;
|
||||
|
|
@ -292,7 +278,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');
|
||||
|
|
@ -309,10 +295,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;
|
||||
|
|
@ -347,85 +333,146 @@ function renderSubtitle(sentence) {
|
|||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
}
|
||||
|
||||
// ===== 渲染说话人气泡 =====
|
||||
// 未确认片段独立展示,不能临时塞进上一位说话人的气泡。
|
||||
// ===== Render: Speaker Bubble =====
|
||||
function renderBubble(sentence) {
|
||||
const id = 'sent-' + sentence.sentence_id;
|
||||
const speakerId = Number(sentence.speaker_id);
|
||||
const trusted = Number.isInteger(speakerId) && speakerId >= 0
|
||||
&& ['fresh', 'confirmed'].includes(sentence.speaker_evidence);
|
||||
const speakerId = sentence.speaker_id;
|
||||
const isUnknown = speakerId < 0;
|
||||
const isInterim = sentence.sentence_type === 0;
|
||||
|
||||
if (isUnknown) {
|
||||
const pendingText = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
let pending = pendingSpanMap[id];
|
||||
if (pending) {
|
||||
pending.spanEl.textContent = pendingText;
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
return;
|
||||
}
|
||||
if (lastConfirmedBubbleEl) {
|
||||
const body = lastConfirmedBubbleEl.querySelector('.bubble-body');
|
||||
const span = document.createElement('span');
|
||||
span.className = 'pending-text';
|
||||
span.dataset.sentenceId = id;
|
||||
span.textContent = pendingText;
|
||||
body.appendChild(span);
|
||||
pendingSpanMap[id] = { hostEl: lastConfirmedBubbleEl, spanEl: span };
|
||||
} else {
|
||||
renderFallbackPendingBubble(id, pendingText, isInterim);
|
||||
}
|
||||
elResultArea.scrollTop = elResultArea.scrollHeight;
|
||||
return;
|
||||
}
|
||||
|
||||
if (pendingSpanMap[id]) {
|
||||
pendingSpanMap[id].spanEl.remove();
|
||||
delete pendingSpanMap[id];
|
||||
}
|
||||
|
||||
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';
|
||||
for (const name of ['speaker-badge', 'speaker-name', 'bubble-time']) {
|
||||
const span = document.createElement('span');
|
||||
span.className = name;
|
||||
header.appendChild(span);
|
||||
}
|
||||
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.append(header, body);
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
el.appendChild(wrapper);
|
||||
elResultArea.appendChild(el);
|
||||
entry = { el };
|
||||
|
||||
if (insertBefore) elResultArea.insertBefore(el, insertBefore);
|
||||
else elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: speakerId };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
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 timeSpan = el.querySelector('.bubble-time');
|
||||
const body = el.querySelector('.bubble-body');
|
||||
body.textContent = sentence.sentence + (isInterim ? ' ...' : '');
|
||||
timeSpan.textContent = formatTimeRange(sentence.start_time, sentence.end_time);
|
||||
body.querySelectorAll('.pending-text').forEach(s => s.remove());
|
||||
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;
|
||||
lastConfirmedBubbleEl = el;
|
||||
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 = '说话人分离已关闭';
|
||||
function renderFallbackPendingBubble(id, text, isInterim) {
|
||||
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);
|
||||
const body = document.createElement('div');
|
||||
body.className = 'bubble-body';
|
||||
wrapper.appendChild(header);
|
||||
wrapper.appendChild(body);
|
||||
el.appendChild(wrapper);
|
||||
elResultArea.appendChild(el);
|
||||
entry = { el: el, speakerId: -1 };
|
||||
sentenceMap[id] = entry;
|
||||
}
|
||||
const body = entry.el.querySelector('.bubble-body');
|
||||
body.textContent = text;
|
||||
body.className = 'bubble-body' + (isInterim ? ' interim' : '');
|
||||
}
|
||||
|
||||
// ===== 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 (!STREAMABLE_AUDIO_EXTENSIONS.has(ext)) {
|
||||
showToast('实时流式测试只支持 PCM 或 WAV,请转换后再试', true);
|
||||
if (UNSUPPORTED_STREAMING_EXTENSIONS.has(ext)) {
|
||||
showToast('当前 demo 不支持直接流式上传 M4A,请转成 WAV 或 MP3', true);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
|
@ -435,9 +482,8 @@ async function startRecognition() {
|
|||
sentenceMap = {};
|
||||
speakerOrderMap = {};
|
||||
speakerOrderCounter = 0;
|
||||
displayStateSupported = false;
|
||||
displayRevision = -1;
|
||||
elSpeakerStatus.textContent = '正在检查说话人服务…';
|
||||
lastConfirmedBubbleEl = null;
|
||||
pendingSpanMap = {};
|
||||
audioChunks = [];
|
||||
elBtnExportWav.disabled = true;
|
||||
elResultPlaceholder?.remove();
|
||||
|
|
@ -450,10 +496,9 @@ async function startRecognition() {
|
|||
sending = true;
|
||||
|
||||
const currentSession = ++sessionId;
|
||||
const useSpeaker = true; // Speaker labels are mandatory for this deployment.
|
||||
let receivedTerminal = false;
|
||||
const useSpeaker = elSpeakerDiarization.checked;
|
||||
|
||||
// 构造 WebSocket 首条 start 消息。
|
||||
// 构造 start 消息
|
||||
let voiceFormat = 0, fileName = '', speedFactor = 0;
|
||||
if (inputMode === 'file') {
|
||||
const ext = getFileExt(selectedFile.name);
|
||||
|
|
@ -464,9 +509,7 @@ async function startRecognition() {
|
|||
|
||||
const startPayload = {
|
||||
type: 'start',
|
||||
model: elEngineModel.value,
|
||||
model_service_url: elModelServiceUrl.value,
|
||||
display_merge: elDisplayMerge.checked,
|
||||
engine_model_type: elEngineModel.value,
|
||||
speaker_diarization: useSpeaker ? 1 : 0,
|
||||
sentence_strategy: parseInt(elSentenceStrategy.value),
|
||||
source: inputMode,
|
||||
|
|
@ -475,23 +518,23 @@ async function startRecognition() {
|
|||
speed_factor: speedFactor
|
||||
};
|
||||
|
||||
const backendUrl = (window.ASR_BACKEND_URL || location.origin).replace(/\/$/, '');
|
||||
const websocketUrl = new URL(backendUrl + '/ws');
|
||||
websocketUrl.protocol = backendUrl.startsWith('https:') ? 'wss:' : 'ws:';
|
||||
ws = new WebSocket(websocketUrl);
|
||||
const protocol = location.protocol === 'https:' ? 'wss:' : 'ws:';
|
||||
ws = new WebSocket(`${protocol}//${location.host}/ws`);
|
||||
ws.binaryType = 'arraybuffer';
|
||||
|
||||
ws.onopen = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
ws.send(JSON.stringify(startPayload));
|
||||
// 在连接尚未建立时点击停止,也要在 start 后补发停止信号。
|
||||
if (!sending) ws.send(JSON.stringify({ type: 'stop' }));
|
||||
if (inputMode === 'file') {
|
||||
sendAudioFile(selectedFile);
|
||||
} else {
|
||||
startMicCapture();
|
||||
}
|
||||
};
|
||||
|
||||
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);
|
||||
}
|
||||
|
|
@ -510,7 +553,8 @@ async function startRecognition() {
|
|||
ws.onclose = () => {
|
||||
if (currentSession !== sessionId) return;
|
||||
stopMicCapture();
|
||||
if (!receivedTerminal) setStatus('error', '连接中断,最终识别结果可能尚未完成');
|
||||
if (sending) setStatus('error', '连接意外断开');
|
||||
else if (stoppingByUser) setStatus('done', '已停止');
|
||||
if (audioChunks.length > 0) elBtnExportWav.disabled = false;
|
||||
ws = null;
|
||||
resetControls();
|
||||
|
|
@ -527,28 +571,10 @@ 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);
|
||||
|
|
@ -557,14 +583,7 @@ 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;
|
||||
|
|
@ -582,7 +601,7 @@ function handleServerMessage(msg, useSpeaker) {
|
|||
}
|
||||
}
|
||||
|
||||
// ===== 停止识别 =====
|
||||
// ===== Stop =====
|
||||
elBtnStop.addEventListener('click', () => stopRecognition());
|
||||
|
||||
function stopRecognition() {
|
||||
|
|
@ -592,12 +611,26 @@ 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) {}
|
||||
}
|
||||
|
||||
// 等待服务端排空 ASR/声纹队列后发送 end,不能用五秒计时器截断更新。
|
||||
|
||||
const stopSession = sessionId;
|
||||
setTimeout(() => {
|
||||
if (sessionId !== stopSession) return;
|
||||
if (ws) {
|
||||
try { ws.close(); } catch (e) {}
|
||||
ws = null;
|
||||
setStatus('done', '已停止(超时)');
|
||||
resetControls();
|
||||
}
|
||||
}, 5000);
|
||||
}
|
||||
|
||||
function resetControls() {
|
||||
|
|
@ -612,53 +645,36 @@ function resetControls() {
|
|||
elBtnStop.disabled = true;
|
||||
}
|
||||
|
||||
// ===== 发送音频文件 =====
|
||||
// 按 16KB 切片发送,并按照音频实际时长等待,确保文件模式也是真实的
|
||||
// 实时输入,而不是瞬间上传完整文件后再由服务端批量切片。
|
||||
// ===== Send Audio File =====
|
||||
// 按 16KB 切片发送,后端会缓冲成 6400 字节块并按 speed_factor 限流
|
||||
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');
|
||||
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) {
|
||||
while (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;
|
||||
const chunkDurationMs = (chunk.byteLength / bytesPerSecond) * 1000 / speedFactor;
|
||||
await new Promise(r => setTimeout(r, Math.max(0, Math.round(chunkDurationMs))));
|
||||
await new Promise(r => setTimeout(r, 0));
|
||||
}
|
||||
if (ownerSession === sessionId && ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
sending = false;
|
||||
setStatus('running', '音频已发送,等待最终结果…');
|
||||
if (ws && ws.readyState === WebSocket.OPEN && sending) {
|
||||
ws.send(JSON.stringify({ type: 'eof' }));
|
||||
}
|
||||
}
|
||||
|
||||
// ===== 麦克风采集 =====
|
||||
// ===== Microphone Capture =====
|
||||
async function startMicCapture() {
|
||||
const ownerSession = sessionId;
|
||||
let stream;
|
||||
try {
|
||||
stream = await navigator.mediaDevices.getUserMedia({
|
||||
micStream = 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;
|
||||
|
|
@ -666,11 +682,6 @@ 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);
|
||||
|
|
@ -725,23 +736,3 @@ function stopMicCapture() {
|
|||
elMicStatus.textContent = '点击下方按钮开始录音';
|
||||
elMicElapsed.textContent = '00:00';
|
||||
}
|
||||
|
||||
// 展示实际部署端点及模型,避免沿用旧 SDK 的无效引擎配置。
|
||||
// The static frontend reads API and WS addresses from runtime-config.js.
|
||||
const backendUrl = (window.ASR_BACKEND_URL || location.origin).replace(/\/$/, '');
|
||||
fetch(backendUrl + '/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,28 +20,22 @@
|
|||
<div class="form-stack">
|
||||
<div class="form-group">
|
||||
<label for="engineModel">引擎模型</label>
|
||||
<input type="text" id="engineModel" value="paraformer-zh-streaming">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="modelServiceUrl">FunASR 本地引擎</label>
|
||||
<input type="text" id="modelServiceUrl" value="local://funasr" readonly>
|
||||
<input type="text" id="engineModel" value="16k_zh_en_speaker" readonly>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label for="sentenceStrategy">分句策略</label>
|
||||
<select id="sentenceStrategy">
|
||||
<option value="0" selected>短停顿(800ms)</option>
|
||||
<option value="1">长停顿(1400ms)</option>
|
||||
<option value="0" selected>语义单句</option>
|
||||
<option value="1">段落</option>
|
||||
</select>
|
||||
</div>
|
||||
<!-- 展示合并与内部音频切段分开配置,便于对照原始片段。 -->
|
||||
<label><input type="checkbox" id="displayMerge" checked> 合并相邻且已确认的同一说话人</label>
|
||||
<div class="form-row">
|
||||
<div class="form-group">
|
||||
<label>话者分离</label>
|
||||
<label class="toggle">
|
||||
<input type="checkbox" id="speakerDiarization" checked disabled>
|
||||
<input type="checkbox" id="speakerDiarization" checked>
|
||||
<span class="toggle-slider"></span>
|
||||
<span class="toggle-label" id="diarizationLabel">必需</span>
|
||||
<span class="toggle-label" id="diarizationLabel">开启</span>
|
||||
</label>
|
||||
</div>
|
||||
</div>
|
||||
|
|
@ -60,9 +54,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" hidden>
|
||||
<input type="file" id="audioFile" accept=".pcm,.wav,.mp3,.aac,.opus,.ogg,.silk,.speex" hidden>
|
||||
</div>
|
||||
<span class="file-info" id="fileInfo">支持 16kHz、单声道、PCM16 的 PCM/WAV</span>
|
||||
<span class="file-info" id="fileInfo">支持 pcm/wav/mp3/aac/opus/silk/speex,m4a 请先转 wav/mp3</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>
|
||||
|
|
@ -85,7 +79,6 @@
|
|||
<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>
|
||||
|
|
@ -113,8 +106,6 @@
|
|||
</main>
|
||||
</div>
|
||||
|
||||
<!-- 版本号变更用于刷新浏览器缓存,确保加载“未知说话人”标签逻辑。 -->
|
||||
<script src="runtime-config.js"></script>
|
||||
<script src="app.js?v=215"></script>
|
||||
<script src="app.js?v=210"></script>
|
||||
</body>
|
||||
</html>
|
||||
|
|
|
|||
|
|
@ -801,5 +801,3 @@ input[type="text"][readonly]:focus {
|
|||
.log-type-end { color: #f9e2af; }
|
||||
.log-type-error { color: #f38ba8; }
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,111 +1,126 @@
|
|||
"""WebSocket contract tests for the FunASR browser adapter."""
|
||||
"""Contract test for the unchanged Tencent UI to FunASR native WS bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
import json
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from aiohttp import web
|
||||
from aiohttp.test_utils import AioHTTPTestCase
|
||||
|
||||
from realtime_websocket.funasr_engine import FunASRSegment
|
||||
from realtime_websocket.funasr_server import (
|
||||
AUXILIARY_SERVICE_KEY,
|
||||
MODEL_SERVICE_KEY,
|
||||
config_handler,
|
||||
AUXILIARY_KEY,
|
||||
SESSION_REGISTRY_KEY,
|
||||
websocket_handler,
|
||||
)
|
||||
|
||||
|
||||
class FakeFunASRSession:
|
||||
class FakeNativeWebSocket:
|
||||
"""Stand in for FunASR's native WSS process without loading model weights."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.sent = False
|
||||
self.incoming: asyncio.Queue[str] = asyncio.Queue()
|
||||
self.audio_bytes = 0
|
||||
self.config: dict[str, object] = {}
|
||||
|
||||
async def feed(self, audio: bytes):
|
||||
if self.sent:
|
||||
return []
|
||||
self.sent = True
|
||||
return [
|
||||
FunASRSegment(
|
||||
text="实时片段",
|
||||
start_time_ms=0,
|
||||
end_time_ms=len(audio) / 32,
|
||||
audio=b"",
|
||||
voiced_ms=len(audio) / 32,
|
||||
is_final=False,
|
||||
sentence_id=0,
|
||||
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})
|
||||
)
|
||||
]
|
||||
|
||||
async def finish(self):
|
||||
return [
|
||||
FunASRSegment(
|
||||
text="最终片段",
|
||||
start_time_ms=0,
|
||||
end_time_ms=1000,
|
||||
audio=b"\x01\x00" * 8000,
|
||||
voiced_ms=1000,
|
||||
is_final=True,
|
||||
sentence_id=0,
|
||||
reason="eof",
|
||||
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})
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
class FakeFunASRService:
|
||||
config = SimpleNamespace(model="fake-funasr")
|
||||
|
||||
def create_session(self):
|
||||
return FakeFunASRSession()
|
||||
async def recv(self) -> str:
|
||||
return await self.incoming.get()
|
||||
|
||||
|
||||
class FakeAuxiliaryService:
|
||||
config = SimpleNamespace(base_url="http://fake-speaker")
|
||||
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 health(self):
|
||||
return {"ready": True, "speaker_embedding_ready": True}
|
||||
|
||||
async def reset_speaker_session(self, session_id):
|
||||
async def reset_speaker_session(self, _session_id):
|
||||
return None
|
||||
|
||||
|
||||
class FunASRWebSocketTests(AioHTTPTestCase):
|
||||
class FunASRBridgeTests(AioHTTPTestCase):
|
||||
def get_app(self):
|
||||
app = web.Application()
|
||||
app[MODEL_SERVICE_KEY] = FakeFunASRService()
|
||||
app[AUXILIARY_SERVICE_KEY] = FakeAuxiliaryService()
|
||||
app.router.add_get("/api/config", config_handler)
|
||||
app[AUXILIARY_KEY] = FakeAuxiliaryService()
|
||||
app[SESSION_REGISTRY_KEY] = {}
|
||||
app.router.add_get("/ws", websocket_handler)
|
||||
return app
|
||||
|
||||
async def test_frontend_contract_uses_funasr_streaming_mode(self):
|
||||
ws = await self.client.ws_connect("/ws")
|
||||
await ws.send_json({"type": "start", "speaker_diarization": 0})
|
||||
start = await ws.receive_json()
|
||||
self.assertEqual(start["type"], "start")
|
||||
self.assertEqual(start["engine"], "funasr")
|
||||
self.assertEqual(start["partial_mode"], "funasr_streaming_cache")
|
||||
async def test_tencent_ui_messages_use_native_funasr_and_keep_speaker_label(self):
|
||||
native = FakeNativeWebSocket()
|
||||
with patch(
|
||||
"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")
|
||||
|
||||
await ws.send_bytes(b"\x01\x00" * 16000)
|
||||
await ws.send_json({"type": "eof"})
|
||||
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
|
||||
messages = []
|
||||
async with asyncio.timeout(5):
|
||||
while True:
|
||||
message = await ws.receive_json()
|
||||
messages.append(message)
|
||||
if message["type"] == "end":
|
||||
break
|
||||
|
||||
sentence_events = [
|
||||
item for item in messages
|
||||
if item["type"] == "sentences" and item["sentences"]
|
||||
]
|
||||
self.assertTrue(any(item["sentences"][0]["sentence_type"] == 0 for item in sentence_events))
|
||||
self.assertEqual(sentence_events[-1]["sentences"][0]["sentence"], "最终片段")
|
||||
self.assertEqual(messages[-1]["type"], "end")
|
||||
await ws.close()
|
||||
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__":
|
||||
|
|
|
|||
|
|
@ -6,3 +6,4 @@ modelscope[framework]==1.34.0
|
|||
soundfile==0.13.1
|
||||
librosa==0.11.0
|
||||
numpy>=1.24
|
||||
websockets>=12,<14
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Start the local CAM++ model service and FunASR WebSocket backend together."""
|
||||
"""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
|
||||
|
|
@ -30,7 +31,7 @@ load_dotenv(PROJECT_ROOT / ".env")
|
|||
|
||||
|
||||
def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
||||
"""Resolve ASR/VAD IDs through the shared model manifest."""
|
||||
"""Resolve a local ASR or VAD model through the project manifest."""
|
||||
name = requested.strip()
|
||||
configured_path = Path(name)
|
||||
direct_candidates = (
|
||||
|
|
@ -56,7 +57,7 @@ def local_model(requested: str, models_dir: Path, kind: str) -> Path:
|
|||
|
||||
|
||||
def local_cam_model(models_dir: Path) -> Path:
|
||||
"""Require one complete CAM++ speaker verification asset from the manifest."""
|
||||
"""Require a complete CAM++ speaker verification asset from the manifest."""
|
||||
manifest = load_manifest()
|
||||
override = os.getenv("CAM_MODEL_PATH", "").strip()
|
||||
if override:
|
||||
|
|
@ -86,7 +87,7 @@ def local_cam_model(models_dir: Path) -> Path:
|
|||
|
||||
|
||||
def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int = 300) -> None:
|
||||
"""Wait until a child has loaded its models, failing when it exits early."""
|
||||
"""Wait for an HTTP child to become ready, reporting early process exit."""
|
||||
deadline = time.monotonic() + seconds
|
||||
while time.monotonic() < deadline:
|
||||
code = process.poll()
|
||||
|
|
@ -103,8 +104,29 @@ def wait_for_health(url: str, process: subprocess.Popen, key: str, seconds: int
|
|||
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 shutdown."""
|
||||
"""Stop a supervised model or WebSocket process on launcher shutdown."""
|
||||
if process is None or process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
|
|
@ -116,7 +138,7 @@ def stop_child(process: subprocess.Popen | None) -> None:
|
|||
|
||||
|
||||
def main() -> None:
|
||||
"""Require ASR, VAD, and CAM++ before exposing the backend WebSocket."""
|
||||
"""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
|
||||
|
|
@ -126,35 +148,102 @@ def main() -> None:
|
|||
)
|
||||
vad = local_model(os.getenv("FUNASR_VAD_MODEL", "fsmn-vad"), models_dir, "vad")
|
||||
cam = local_cam_model(models_dir)
|
||||
print(f"Local models: ASR={asr}; VAD={vad}; CAM++={cam}", flush=True)
|
||||
|
||||
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()
|
||||
env.update({
|
||||
"MODEL_DIR": str(models_dir),
|
||||
"FUNASR_ASR_MODEL": str(asr),
|
||||
"FUNASR_VAD_MODEL": str(vad),
|
||||
"CAM_MODEL_PATH": str(cam),
|
||||
"AUXILIARY_SERVICE_URL": "http://127.0.0.1:8010",
|
||||
})
|
||||
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", "scripts.auxiliary_server"],
|
||||
cwd=PROJECT_ROOT, env=env,
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
wait_for_health("http://127.0.0.1:8010/health", auxiliary, "speaker_embedding_ready")
|
||||
wait_for_health(f"{aux_url}/health", auxiliary, "speaker_embedding_ready")
|
||||
|
||||
native_args = [
|
||||
sys.executable,
|
||||
str(PROJECT_ROOT / "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", "scripts.run_funasr_demo", "--no-browser"],
|
||||
cwd=PROJECT_ROOT, env=env,
|
||||
cwd=PROJECT_ROOT,
|
||||
env=env,
|
||||
)
|
||||
web_port = int(env.get("WEB_PORT", "8082"))
|
||||
wait_for_health(
|
||||
f"http://127.0.0.1:{web_port}/api/config", websocket, "engine"
|
||||
)
|
||||
print(f"Backend ready: ws://127.0.0.1:{web_port}/ws", flush=True)
|
||||
print(f"FunASR browser backend ready: ws://127.0.0.1:{web_port}/ws", flush=True)
|
||||
while True:
|
||||
for label, process in (("CAM++", auxiliary), ("WebSocket", websocket)):
|
||||
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})")
|
||||
|
|
@ -163,6 +252,7 @@ def main() -> None:
|
|||
pass
|
||||
finally:
|
||||
stop_child(websocket)
|
||||
stop_child(native)
|
||||
stop_child(auxiliary)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,56 +1,131 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Serve the browser UI independently of the FunASR backend."""
|
||||
"""Serve the unchanged Tencent demo UI and proxy its API to the backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import asyncio
|
||||
import os
|
||||
from http.server import SimpleHTTPRequestHandler, ThreadingHTTPServer
|
||||
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 = PROJECT_ROOT / "realtime_websocket" / "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)
|
||||
|
||||
|
||||
class FrontendHandler(SimpleHTTPRequestHandler):
|
||||
"""Serve static assets and a browser-visible backend URL at runtime."""
|
||||
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}"
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, directory=str(STATIC_ROOT), **kwargs)
|
||||
|
||||
def do_GET(self) -> None:
|
||||
if self.path.split("?", 1)[0] != "/runtime-config.js":
|
||||
return super().do_GET()
|
||||
backend_url = os.getenv("BACKEND_PUBLIC_URL", "http://127.0.0.1:8082").rstrip("/")
|
||||
parsed = urlsplit(backend_url)
|
||||
if parsed.scheme not in {"http", "https"} or not parsed.netloc or parsed.path:
|
||||
self.send_error(500, "BACKEND_PUBLIC_URL must be an HTTP origin")
|
||||
return
|
||||
# JSON escaping also yields a valid JavaScript string literal.
|
||||
body = ("window.ASR_BACKEND_URL = " + json.dumps(backend_url) + ";\n").encode("utf-8")
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "application/javascript; charset=utf-8")
|
||||
self.send_header("Cache-Control", "no-store")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
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:
|
||||
host = os.getenv("FRONTEND_HOST", "127.0.0.1")
|
||||
port = int(os.getenv("FRONTEND_PORT", "8080"))
|
||||
server = ThreadingHTTPServer((host, port), FrontendHandler)
|
||||
print(f"Frontend: http://{host}:{port}/", flush=True)
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
finally:
|
||||
server.server_close()
|
||||
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__":
|
||||
|
|
|
|||
Loading…
Reference in New Issue