ASR-demo/backend/run_backend.py

275 lines
9.7 KiB
Python

#!/usr/bin/env python3
"""启动本地 CAM++、FunASR 原生实时 WSS 服务和浏览器协议桥接层。"""
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]
# 直接从 backend/ 执行脚本时,该目录会排在搜索路径首位;
# 将仓库根目录提前加入,确保后端模块和本地模型清单都能正确导入。
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:
"""通过项目模型清单查找本地 ASR 或 VAD 模型。"""
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:
"""从模型清单中确认 CAM++ 说话人验证资源完整。"""
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:
"""等待 HTTP 子进程就绪;若进程提前退出则报告错误。"""
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:
"""等待 FunASR 加载模型并打开内部 WebSocket。"""
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
# 遇到端口占用时明确报错,避免误将其他进程的端口视为本服务。
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:
"""启动器关闭时停止受其管理的模型或 WebSocket 进程。"""
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:
"""确认 ASR、VAD 和 CAM++ 均可用后,再开放公网 WebSocket 桥接服务。"""
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"
)
hotword_asr = local_model(
os.getenv("FUNASR_HOTWORD_ASR_MODEL", "paraformer-zh-contextual"),
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: streaming ASR={asr}; hotword final ASR={hotword_asr}; "
f"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()
}
# 实时 VAD 由 FunASR 原生服务负责;不要在 CAM++ 进程中重复加载。
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_HOTWORD_ASR_MODEL": str(hotword_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:
# 先加载 ASR 主模型,避免说话人辅助模型抢占显存导致识别模型无法启动。
native_args = [
sys.executable,
str(PROJECT_ROOT / "backend" / "realtime_websocket" / "funasr_native_wss.py"),
"--host",
native_host,
"--port",
str(native_port),
"--asr_model",
str(hotword_asr),
"--asr_model_revision",
os.getenv("FUNASR_HOTWORD_ASR_MODEL_REVISION", "v2.0.4"),
"--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)
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")
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 (
("FunASR native WS", native),
("CAM++", auxiliary),
("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()