ASR-demo/scripts/serve.py

217 lines
8.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

#!/usr/bin/env python3
"""在宿主机启动独立的 Qwen3-ASR VLLM 服务。"""
from __future__ import annotations
import argparse
import os
import shutil
import signal
import subprocess
import sys
import time
from pathlib import Path
from urllib.error import URLError
from urllib.request import urlopen
from dotenv import load_dotenv
# 启动器自动读取 demo/.env系统环境变量仍然优先便于部署平台临时覆盖配置。
PROJECT_ROOT = Path(__file__).resolve().parents[1]
load_dotenv(PROJECT_ROOT / ".env")
# 将服务端口集中在代码变量中维护,启动时不需要额外传入端口参数;健康检查、
# VLLM 子进程命令和就绪提示都使用同一个端口,避免配置不一致导致误判。
SERVER_PORT = int(os.getenv("VLLM_PORT", "9950"))
try:
from .model_manifest import load_manifest, model_directory, resolve_model_id
except ImportError:
from model_manifest import load_manifest, model_directory, resolve_model_id
def has_model_weights(model_path: Path) -> bool:
"""检查 VLLM 加载模型前必须存在的最小本地文件集合。"""
if not model_path.is_dir() or not (model_path / "config.json").is_file():
return False
return any(model_path.rglob("*.safetensors")) or any(model_path.rglob("*.bin"))
def positive_int(value: str) -> int:
"""解析启动轮询使用的正整数参数,并拒绝零和负数。"""
parsed = int(value)
if parsed < 1:
raise argparse.ArgumentTypeError("value must be at least 1")
return parsed
def non_negative_float(value: str) -> float:
"""解析启动轮询间隔,并拒绝会导致逻辑异常的负数。"""
parsed = float(value)
if parsed < 0:
raise argparse.ArgumentTypeError("value must be non-negative")
return parsed
def build_parser() -> argparse.ArgumentParser:
"""创建宿主机启动参数解析器,默认值允许通过环境变量统一覆盖。"""
parser = argparse.ArgumentParser(description=__doc__)
add_arguments(parser)
return parser
def add_arguments(parser: argparse.ArgumentParser) -> None:
"""注册模型、网络端点和启动检查循环相关的命令行参数。"""
parser.add_argument(
"--model",
default=os.getenv("QWEN3_ASR_MODEL", "default"),
help="Model alias, exact model ID, or default",
)
parser.add_argument(
"--models-dir",
type=Path,
default=Path(os.getenv("MODEL_DIR", str(Path(__file__).resolve().parents[1] / "models"))),
help="Root directory containing downloaded model files",
)
parser.add_argument("--host", default=os.getenv("VLLM_HOST", "0.0.0.0"))
parser.add_argument(
"--display-host",
default=os.getenv("VLLM_DISPLAY_HOST", "127.0.0.1"),
help="Host name shown in the ready message; does not change the bind address",
)
parser.add_argument(
"--probe-host",
default=os.getenv("VLLM_PROBE_HOST", "127.0.0.1"),
help="Host used by the startup health probe",
)
parser.add_argument(
"--startup-check-loops",
type=positive_int,
default=positive_int(os.getenv("VLLM_STARTUP_CHECK_LOOPS", "60")),
help="Maximum number of health checks before startup fails",
)
parser.add_argument(
"--startup-check-interval",
type=non_negative_float,
default=non_negative_float(os.getenv("VLLM_STARTUP_CHECK_INTERVAL_SECONDS", "2")),
help="Seconds between startup health checks",
)
parser.add_argument("--served-model-name", default=os.getenv("VLLM_SERVED_MODEL_NAME"))
parser.add_argument(
"--gpu-memory-utilization",
default=os.getenv("VLLM_GPU_MEMORY_UTILIZATION", "0.3"),
)
parser.add_argument("--max-model-len", default=os.getenv("VLLM_MAX_MODEL_LEN", "16384"))
parser.add_argument("--max-num-seqs", default=os.getenv("VLLM_MAX_NUM_SEQS", "16"))
parser.add_argument("--tensor-parallel-size", default=os.getenv("VLLM_TENSOR_PARALLEL_SIZE", "1"))
parser.add_argument(
"--enforce-eager",
action=argparse.BooleanOptionalAction,
default=os.getenv("VLLM_ENFORCE_EAGER", "true").lower() == "true",
)
def build_server_command(args: argparse.Namespace, model_id: str, model_path: Path) -> list[str]:
"""构造新版 VLLM 原生启动命令,不依赖 qwen-asr-serve。"""
# Qwen3-ASR 已由新版 VLLM 原生支持,因此这里调用 vllm serve避免
# qwen-asr-serve 对旧版 VLLM 的固定依赖影响 GB10 部署环境。
executable_name = os.getenv("VLLM_EXECUTABLE", "vllm")
executable = shutil.which(executable_name)
if executable is None:
raise RuntimeError(f"{executable_name} was not found; install requirements-deploy.txt first")
served_model_name = args.served_model_name or model_id
command = [
executable,
"serve",
str(model_path),
"--host",
args.host,
"--port",
str(SERVER_PORT),
"--served-model-name",
served_model_name,
"--gpu-memory-utilization",
str(args.gpu_memory_utilization),
"--max-model-len",
str(args.max_model_len),
"--max-num-seqs",
str(args.max_num_seqs),
"--tensor-parallel-size",
str(args.tensor_parallel_size),
]
if args.enforce_eager:
command.append("--enforce-eager")
return command
def wait_until_ready(process: subprocess.Popen[bytes], probe_url: str, loops: int, interval: float) -> None:
"""按调用方指定的次数和间隔轮询 VLLM 健康接口,直到服务就绪或失败。"""
for attempt in range(1, loops + 1):
if process.poll() is not None:
raise RuntimeError(f"VLLM exited during startup with code {process.returncode}")
try:
with urlopen(probe_url, timeout=2) as response:
if 200 <= response.status < 300:
return
except (OSError, URLError):
pass
print(f"Waiting for VLLM startup ({attempt}/{loops})...", flush=True)
if attempt < loops:
time.sleep(interval)
raise TimeoutError(f"VLLM did not become ready after {loops} health checks: {probe_url}")
def stop_process(process: subprocess.Popen[bytes]) -> None:
"""向 VLLM 子进程转发优雅停止信号,并在超时后执行兜底清理。"""
if process.poll() is not None:
return
if os.name == "nt":
process.send_signal(signal.CTRL_BREAK_EVENT)
else:
process.send_signal(signal.SIGINT)
try:
process.wait(timeout=10)
except subprocess.TimeoutExpired:
process.terminate()
process.wait(timeout=10)
def main() -> int:
"""解析一个本地模型、启动 VLLM并持续托管宿主机子进程。"""
args = build_parser().parse_args()
manifest = load_manifest()
model_id = resolve_model_id(args.model, manifest)
model_path = model_directory(model_id, manifest, args.models_dir.resolve())
if not has_model_weights(model_path):
print(f"Model is missing or incomplete: {model_id} ({model_path})", file=sys.stderr)
print("Run scripts/download_models.py for the same model first.", file=sys.stderr)
return 1
command = build_server_command(args, model_id, model_path)
probe_url = f"http://{args.probe_host}:{SERVER_PORT}/health"
display_url = f"http://{args.display_host}:{SERVER_PORT}"
print(f"Starting Qwen3-ASR VLLM service on host: {display_url}", flush=True)
print(f"Model: {model_id}", flush=True)
print(f"Startup checks: {args.startup_check_loops} x {args.startup_check_interval}s", flush=True)
process = subprocess.Popen(command)
try:
wait_until_ready(process, probe_url, args.startup_check_loops, args.startup_check_interval)
print(f"VLLM ready: {display_url}", flush=True)
print(f"OpenAI endpoint: {display_url}/v1", flush=True)
while process.poll() is None:
time.sleep(0.5)
return int(process.returncode or 0)
except (KeyboardInterrupt, TimeoutError, RuntimeError) as exc:
print(str(exc), file=sys.stderr)
return 1
finally:
stop_process(process)
if __name__ == "__main__":
raise SystemExit(main())