ASR-demo/realtime_asr_optimization_demo/server.py

912 lines
44 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.

"""面向浏览器的独立 Qwen3-ASR VLLM WebSocket 编排服务。"""
from __future__ import annotations
import argparse
import asyncio
import json
import math
import logging
import os
import time
import webbrowser
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
from uuid import uuid4
from aiohttp import WSMsgType, web
from dotenv import load_dotenv
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
from model_service import ModelServiceConfig, VLLMTranscriptionService
from speaker_assembler import SegmentAssembler
# 与部署启动器读取同一配置;外部环境变量优先于 demo/.env。
DEPLOY_ROOT = Path(__file__).resolve().parents[1]
load_dotenv(DEPLOY_ROOT / ".env")
# 监听所有网卡,允许同一局域网内的浏览器访问服务器上的 Demo端口集中在代码
# 变量中维护,便于服务器部署时直接修改并保持页面和 WebSocket 使用一致端口。
WEB_HOST = "0.0.0.0"
WEB_PORT = 8082
WEB_DISPLAY_HOST = os.getenv("WEB_DISPLAY_HOST", "127.0.0.1")
DEFAULT_MODEL_SERVICE_URL = f"http://127.0.0.1:{os.getenv('VLLM_PORT', '9950')}/v1"
PARTIAL_BYTES_PER_SECOND = 16000 * 2
VAD_FRAME_BYTES = 640
VAD_FRAME_MS = 20
VAD_SILENCE_MS = 800
PARAGRAPH_SILENCE_MS = 1400
VAD_RMS_THRESHOLD = 450
MIN_SPEAKER_VOICE_MS = 800
# 说话人切换感知切段:活跃 turn 内每积累一个窗口就用声纹比对一次,
# 与当前 turn 参考向量的余弦低于阈值时视为换人,在窗口边界强制切段。
# 静音阈值800/1400ms覆盖不了常见的 200~600ms 换话停顿,窗比对是它的补充。
SPEAKER_TURN_WINDOW_MS = max(800, int(os.getenv("SPEAKER_TURN_WINDOW_MS", "1000")))
SPEAKER_SWITCH_SIMILARITY = float(os.getenv("SPEAKER_SWITCH_SIMILARITY", "0.5"))
SPEAKER_WINDOW_MAX_FAILURES = 3
# 窗口里的有声时长不足时比对没有意义(长静音窗声纹不可信),攒够再查,
# 同时避免静音段每帧重复触发请求。
SPEAKER_WINDOW_MIN_VOICED_MS = 600
# 失配窗内定位精确切换点:换话几乎总有停顿,取首个"停顿≥该时长后的起音"
# 作为切点,避免把上一位的话尾按窗口起点粗暴划给下一位。
SPEAKER_SWITCH_GAP_MS = 200
LOGGER = logging.getLogger(__name__)
def _coerce_similarity(value: Any, default: float) -> float:
"""解析 start 消息里的相似度阈值,非法值回退到默认,避免 NaN 进比较。"""
try:
parsed = float(value)
except (TypeError, ValueError):
return default
return parsed if math.isfinite(parsed) and -1.0 <= parsed <= 1.0 else default
class EndOfStream:
"""带明确类型的队列结束标记,用于区分控制信号和真实音频字节。"""
EOF = EndOfStream()
MODEL_SERVICE_KEY = web.AppKey("model_service", VLLMTranscriptionService)
AUXILIARY_SERVICE_KEY = web.AppKey("auxiliary_service", AuxiliaryModelService)
@dataclass(frozen=True)
class SpeakerJob:
"""等待辅助服务处理的单个已完成 turn只保存该 turn 的 PCM 音频。"""
sentence_id: int
audio: bytes
start_time_ms: float
end_time_ms: float
voiced_ms: float = 0.0
# 强制切段产生的头段可能混入下一说话人至多一个窗口的音频,标记后
# 辅助服务只用它匹配标签、不更新簇质心,避免污染在线聚类。
speaker_verified: bool = True
def validate_model_service_url(value: str) -> str:
"""只接受用户输入的 HTTP(S) VLLM 地址,并拒绝附带认证和查询参数的地址。"""
candidate = value.strip().rstrip("/")
parsed = urlparse(candidate)
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError("VLLM 地址必须是完整的 http:// 或 https:// URL")
if parsed.username or parsed.password or parsed.query or parsed.fragment:
raise ValueError("VLLM 地址不能包含账号、密码、查询参数或片段")
return candidate
@dataclass
class SessionMetrics:
"""记录 WebSocket 会话耗时和结果修订次数,并在会话结束后展示。"""
started_at: float
audio_bytes: int = 0
input_chunks: int = 0
partial_count: int = 0
partial_revisions: int = 0
first_partial_ms: float | None = None
final_ms: float | None = None
def snapshot(self) -> dict[str, Any]:
"""返回可安全序列化为 JSON 的指标,耗时均相对于会话开始时间计算。"""
now = time.perf_counter()
return {
"audio_bytes": self.audio_bytes,
"input_chunks": self.input_chunks,
"partial_count": self.partial_count,
"partial_revisions": self.partial_revisions,
"first_partial_ms": self.first_partial_ms,
"final_ms": self.final_ms,
"elapsed_ms": round((now - self.started_at) * 1000, 1),
}
class RealtimeSession:
"""串行处理音频快照,并通过同一个 WebSocket 有序推送状态更新。"""
def __init__(
self,
ws: web.WebSocketResponse,
model_service: VLLMTranscriptionService,
auxiliary_service: AuxiliaryModelService | None,
start: dict[str, Any],
) -> None:
self.ws = ws
self.model_service = model_service
self.auxiliary_service = auxiliary_service
self.start = start
# 聚类状态只能属于当前连接,客户端复用 ID 不能串入另一会话的声纹池。
self.session_id = uuid4().hex
self.send_lock = asyncio.Lock()
self.state_revision = 0
self.audio_queue: asyncio.Queue[bytes | EndOfStream] = asyncio.Queue(maxsize=256)
self.speaker_queue: asyncio.Queue[SpeakerJob | EndOfStream] = asyncio.Queue(maxsize=64)
self.assembler = SegmentAssembler()
self.metrics = SessionMetrics(time.perf_counter())
self.source = str(start.get("source") or "mic")
self.file_name = str(start.get("file_name") or "audio.wav")
self.windowed_partial = self.source == "mic" or Path(self.file_name).suffix.lower() in {".pcm", ".wav"}
self.sentence_strategy = int(start.get("sentence_strategy") or 0)
self.silence_limit_ms = PARAGRAPH_SILENCE_MS if self.sentence_strategy == 1 else VAD_SILENCE_MS
self.partial_interval_ms = max(300, int(start.get("partial_interval_ms") or 1200))
self.max_segment_sec = max(2.0, float(start.get("max_segment_sec") or 12.0))
self.merge_adjacent = self._parse_flag(start.get("display_merge"), True)
self.enable_native_partial = self._parse_flag(start.get("enable_native_partial_stream"), True)
self.segment_id = 0
self.segment_audio = bytearray()
self.segment_start_ms = 0.0
self.vad_buffer = bytearray()
self.processed_audio_bytes = 0
self.silence_ms = 0
self.in_speech = False
self.voiced_ms = 0.0
self.pre_roll = bytearray()
# 窗级说话人切换检测状态:参考向量是当前段已通过比对窗口的均值,
# window_* 记录自上次比对以来累积的音频字节与有声时长。
self.turn_window_ms = max(800, int(start.get("speaker_window_ms") or SPEAKER_TURN_WINDOW_MS))
self.switch_similarity = _coerce_similarity(start.get("speaker_switch_similarity"), SPEAKER_SWITCH_SIMILARITY)
self.turn_window_bytes = int(self.turn_window_ms / 1000 * PARTIAL_BYTES_PER_SECOND)
self.turn_reference: list[float] | None = None
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
self.window_check_failures = 0
self.window_check_available = (
auxiliary_service is not None
and hasattr(auxiliary_service, "speaker_embedding")
)
self.wav_header_buffer = bytearray()
self.wav_payload_started = self.source != "file" or Path(self.file_name).suffix.lower() != ".wav"
self.wav_riff_read = False
self.wav_format_valid = False
self.wav_data_remaining: int | None = None
self.speaker_warning_sent = False
self.speaker_enabled = self._parse_flag(start.get("speaker_diarization"), True)
self.window_check_available = self.window_check_available and self.speaker_enabled
self.input_stopped = False
@staticmethod
def _parse_flag(value: Any, default: bool) -> bool:
"""兼容前端传来的 0/1、布尔值和字符串开关避免字符串 0 被误判为真。"""
if value is None:
return default
if isinstance(value, str):
return value.strip().lower() not in {"", "0", "false", "no", "off"}
return bool(value)
async def emit(self, payload: dict[str, Any]) -> None:
"""在连接仍然有效时发送一条有序事件,避免向已关闭连接写入数据。"""
async with self.send_lock:
if not self.ws.closed:
await self.ws.send_json(payload)
async def emit_state(self, sentence: dict[str, Any] | None = None) -> None:
"""每次状态更新后同时发送原始状态和重新计算的展示快照。"""
# 在首次 await 前冻结快照,音频 worker 与 speaker worker 不会混用两版状态。
self.state_revision += 1
state = {
"type": "display_state", "revision": self.state_revision,
"raw_segments": self.assembler.raw_snapshot(),
"display_blocks": self.assembler.display_blocks(self.merge_adjacent),
"metrics": self.metrics.snapshot(),
}
if sentence is not None:
await self.emit({"type": "sentences", "sentences": [sentence], "metrics": self.metrics.snapshot()})
await self.emit(state)
async def warn_speaker(self, message: str) -> None:
"""只发送一次说话人服务告警,避免辅助服务异常时刷屏。"""
if self.speaker_warning_sent:
return
await self.emit(
{
"type": "speaker_warning",
"session_id": self.session_id,
"speaker_service_url": getattr(getattr(self.auxiliary_service, "config", None), "base_url", None),
"message": message,
}
)
self.speaker_warning_sent = True
def _duration_ms(self) -> float:
"""根据 16 kHz PCM 字节数计算时长,不依赖前端可能漂移的时间戳。"""
return len(self.segment_audio) / PARTIAL_BYTES_PER_SECOND * 1000
async def _transcribe(self, partial: bool, audio: bytes | None = None) -> str | None:
"""通过 VLLM 适配器转写当前逻辑片段;换人切段时转写显式传入的头音频。"""
return await self.model_service.transcribe(
bytes(self.segment_audio) if audio is None else audio,
"mic",
"turn.pcm",
partial=partial,
)
async def _emit_transcription(self, text: str, sentence_type: int, end_ms: float, commit_reason: str | None = None) -> None:
"""写入或更新一条句子,确保中间结果和最终结果不会在前端产生重复行。"""
if not text:
return
sentence = self.assembler.apply_sentence(
{
"sentence_id": self.segment_id,
"sentence": text,
"sentence_type": sentence_type,
"start_time": self.segment_start_ms,
"end_time": end_ms,
"speaker_id": -1,
"speaker_name": "",
"speaker_evidence": "pending",
"speaker_confidence": 0.0,
"speaker_strategy": "vllm_no_speaker_evidence",
"commit_reason": commit_reason,
"speaker_status": ("queued" if sentence_type else "waiting_final") if self.speaker_enabled else "disabled",
"speaker_reason": ("等待声纹处理" if sentence_type else "语音片段结束后识别说话人") if self.speaker_enabled else "说话人分离已关闭",
}
)
if sentence_type == 0:
self.metrics.partial_count += 1
if sentence["revision_count"] > 0:
self.metrics.partial_revisions += 1
if self.metrics.first_partial_ms is None:
self.metrics.first_partial_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
else:
self.metrics.final_ms = round((time.perf_counter() - self.metrics.started_at) * 1000, 1)
await self.emit_state(sentence)
@staticmethod
def _is_voice_frame(frame: bytes) -> bool:
"""用 PCM 帧的 RMS 判断是否有语音,作为实时低延迟切句触发器。"""
if not frame:
return False
samples = memoryview(frame).cast("h")
if not samples:
return False
square_mean = sum(sample * sample for sample in samples) / len(samples)
return math.sqrt(square_mean) >= VAD_RMS_THRESHOLD
def _strip_wav_header(self, chunk: bytes) -> bytes:
"""增量解析 RIFF chunk支持扩展头并拒绝采样率或声道不匹配的 WAV。"""
if self.wav_payload_started:
if self.wav_data_remaining is None:
return chunk
payload = chunk[:self.wav_data_remaining]
self.wav_data_remaining -= len(payload)
return payload
self.wav_header_buffer.extend(chunk)
buffer = self.wav_header_buffer
if not self.wav_riff_read:
if len(buffer) < 12:
return b""
if buffer[:4] != b"RIFF" or buffer[8:12] != b"WAVE":
raise ValueError("文件不是有效的 RIFF/WAV 音频")
del buffer[:12]
self.wav_riff_read = True
while len(buffer) >= 8:
kind = bytes(buffer[:4])
size = int.from_bytes(buffer[4:8], "little")
if kind == b"data":
if not self.wav_format_valid or size % 2:
raise ValueError("WAV 必须为 16kHz、单声道、PCM16")
self.wav_data_remaining = size
self.wav_payload_started = True
payload = bytes(buffer[8:8 + size])
self.wav_data_remaining -= len(payload)
buffer.clear()
return payload
if size > 1024 * 1024:
raise ValueError("WAV 元数据头过大,请转换为标准 PCM WAV")
chunk_size = 8 + size + (size % 2)
if len(buffer) < chunk_size:
return b""
if kind == b"fmt ":
fmt = buffer[8:8 + size]
fields = (int.from_bytes(fmt[0:2], "little"), int.from_bytes(fmt[2:4], "little"),
int.from_bytes(fmt[4:8], "little"), int.from_bytes(fmt[14:16], "little"))
if size < 16 or fields != (1, 1, 16000, 16):
raise ValueError("WAV 必须为 16kHz、单声道、PCM16请先转换音频")
self.wav_format_valid = True
del buffer[:chunk_size]
return b""
async def _resolve_speaker(self, job: SpeakerJob) -> None:
"""异步解析单个 turn 的说话人,并把结果覆盖回同一个 sentence_id。"""
if not self.speaker_enabled:
return
async def update_status(status: str, reason: str) -> None:
"""将每个失败或等待阶段回写原片段,避免只发一次全局告警。"""
updated = self.assembler.apply_speaker_update({
"sentence_id": job.sentence_id, "speaker_id": -1,
"speaker_evidence": "pending", "speaker_confidence": 0.0,
"speaker_status": status, "speaker_reason": reason,
})
await self.emit_state(updated)
# 按有效有声帧检查长度,不能让句尾 800ms 静音把短插话伪装成长样本。
if job.voiced_ms < MIN_SPEAKER_VOICE_MS:
await update_status("insufficient_audio", f"有效语音不足 {MIN_SPEAKER_VOICE_MS}ms不继承上一位说话人")
return
if self.auxiliary_service is None:
await update_status("service_unavailable", "未配置说话人辅助模型服务")
await self.warn_speaker("未配置辅助模型服务,无法执行实时说话人分离")
return
await update_status("processing", "正在提取声纹并匹配说话人")
try:
speaker = await self.auxiliary_service.resolve_speaker(
job.audio,
self.session_id,
job.start_time_ms,
job.end_time_ms,
job.speaker_verified,
)
except Exception as exc:
# 辅助服务异常不能阻断 ASR当前片段继续保持 pending方便定位服务问题。
LOGGER.exception("speaker resolve failed: session=%s sentence=%s", self.session_id, job.sentence_id)
await update_status("service_error", str(exc))
await self.warn_speaker(str(exc))
return
if not speaker:
await update_status("no_embedding", "辅助服务未返回可用声纹结果")
return
update = dict(speaker)
update["sentence_id"] = job.sentence_id
update["speaker_name"] = str(update.get("speaker_name") or "")
updated = self.assembler.apply_speaker_update(update)
if updated is not None:
LOGGER.info("speaker result: session=%s sentence=%s status=%s strategy=%s", self.session_id,
job.sentence_id, updated.get("speaker_status"), updated.get("speaker_strategy"))
await self.emit_state(updated)
async def process_speakers(self) -> None:
"""按 turn 顺序串行访问辅助模型,保证在线聚类中心不会乱序更新。"""
while True:
item = await self.speaker_queue.get()
if isinstance(item, EndOfStream):
return
await self._resolve_speaker(item)
async def _commit_segment(self, reason: str = "final") -> None:
"""在 VAD 检测到一句结束后提交 final并异步排队当前 turn 的说话人解析。"""
if not self.segment_audio or not self.in_speech:
return
# 去掉句尾触发切段的静音ASR 与声纹都使用当前片段的真实有效范围。
trailing_bytes = int(self.silence_ms * PARTIAL_BYTES_PER_SECOND / 1000)
if trailing_bytes:
del self.segment_audio[-trailing_bytes:]
final_audio = bytes(self.segment_audio)
final_start_ms = self.segment_start_ms
final_end_ms = self.segment_start_ms + self._duration_ms()
final_sentence_id = self.segment_id
text = await self._transcribe(partial=False)
if text:
await self._emit_transcription(text, 1, final_end_ms, reason)
if self.speaker_enabled:
await self.speaker_queue.put(
SpeakerJob(
sentence_id=final_sentence_id,
audio=final_audio,
start_time_ms=final_start_ms,
end_time_ms=final_end_ms,
voiced_ms=self.voiced_ms,
)
)
elif self.assembler.segments.pop(final_sentence_id, None) is not None:
# final 判定无文本时撤回临时结果,不能留下永远等待声纹的 partial。
await self.emit_state()
self.segment_audio.clear()
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
self.segment_id += 1
self.silence_ms = 0
self.in_speech = False
self.voiced_ms = 0.0
self._reset_turn_embedding_state()
def _reset_turn_embedding_state(self) -> None:
"""段与段之间不共享声纹参考;新 turn 从第一个窗口重新建立参考。"""
self.turn_reference = None
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
@staticmethod
def _unit_average(first: list[float], second: list[float]) -> list[float]:
"""两个单位向量的等权平均并重新归一化,作为滚动参考向量。"""
merged = [a + b for a, b in zip(first, second)]
norm = math.sqrt(sum(value * value for value in merged))
if not math.isfinite(norm) or norm < 1e-8:
return list(first)
return [value / norm for value in merged]
async def _check_speaker_window(self) -> bool:
"""比对活跃段尾部窗口的声纹,检出换人时在窗口边界强制切段。
窗提取失败绝不阻断 ASR连续失败或辅助服务不支持时整体停用
与既有 speaker 失败不阻断转写的策略一致。返回是否发生了切段。"""
check_len = len(self.segment_audio)
window_audio = bytes(self.segment_audio[self.window_start_bytes:check_len])
embedding = await self.auxiliary_service.speaker_embedding(window_audio)
if embedding is None:
if getattr(self.auxiliary_service, "speaker_embedding_unsupported", False):
self.window_check_available = False
await self.warn_speaker("辅助服务不支持窗口声纹提取,已停用说话人切换感知切段")
else:
self.window_check_failures += 1
if self.window_check_failures >= SPEAKER_WINDOW_MAX_FAILURES:
self.window_check_available = False
await self.warn_speaker("窗口声纹提取连续失败,本次会话已停用切换感知切段")
return False
self.window_check_failures = 0
reference = self.turn_reference
if reference is None or len(reference) != len(embedding):
self.turn_reference = embedding
self.window_start_bytes = check_len
self.window_voiced_ms = 0.0
return False
score = sum(a * b for a, b in zip(reference, embedding))
if score >= self.switch_similarity:
self.turn_reference = self._unit_average(reference, embedding)
self.window_start_bytes = check_len
self.window_voiced_ms = 0.0
return False
split_ms = self.window_start_bytes / PARTIAL_BYTES_PER_SECOND * 1000
LOGGER.info(
"speaker switch detected: session=%s sentence=%s score=%.3f threshold=%.2f at_ms=%.0f",
self.session_id, self.segment_id, score, self.switch_similarity,
self.segment_start_ms + split_ms,
)
await self._split_on_speaker_change(self.window_start_bytes, embedding)
return True
def _find_switch_boundary(self, window_start: int) -> tuple[int, float]:
"""在失配窗内定位精确切换点:首个"停顿≥SPEAKER_SWITCH_GAP_MS 后的起音"
找不到明显停顿就保守回退到窗口起点。返回切点字节偏移与切点之后的
有声时长,供头尾两段正确归属 voiced_ms。"""
gap_frames = max(1, SPEAKER_SWITCH_GAP_MS // VAD_FRAME_MS)
offsets = range(window_start, len(self.segment_audio) - VAD_FRAME_BYTES + 1, VAD_FRAME_BYTES)
voiced_flags = [
self._is_voice_frame(bytes(self.segment_audio[offset:offset + VAD_FRAME_BYTES]))
for offset in offsets
]
boundary = window_start
silence_run = 0
for offset, voiced in zip(offsets, voiced_flags):
if voiced and silence_run >= gap_frames:
boundary = offset
break
silence_run = 0 if voiced else silence_run + 1
tail_voiced_ms = sum(
VAD_FRAME_MS for offset, voiced in zip(offsets, voiced_flags)
if voiced and offset >= boundary
)
return boundary, float(tail_voiced_ms)
async def _split_on_speaker_change(self, window_start: int, new_reference: list[float]) -> None:
"""窗比对检出不同说话人:头段提交 final 并排队声纹,尾段直接成为新 turn。
与 _commit_segment 不同:切点位于活跃语音内部,不做尾部静音裁剪,
尾段保留全部已收音频并沿用被判为新说话人的窗口向量作为新参考。"""
split_at, tail_voiced_ms = self._find_switch_boundary(window_start)
head_audio = bytes(self.segment_audio[:split_at])
head_duration_ms = split_at / PARTIAL_BYTES_PER_SECOND * 1000
head_start_ms = self.segment_start_ms
head_end_ms = head_start_ms + head_duration_ms
head_sentence_id = self.segment_id
head_voiced_ms = max(0.0, self.voiced_ms - tail_voiced_ms)
text = await self._transcribe(partial=False, audio=head_audio)
if text:
await self._emit_transcription(text, 1, head_end_ms, "speaker_change")
if self.speaker_enabled:
await self.speaker_queue.put(
SpeakerJob(
sentence_id=head_sentence_id,
audio=head_audio,
start_time_ms=head_start_ms,
end_time_ms=head_end_ms,
voiced_ms=head_voiced_ms,
speaker_verified=False,
)
)
elif self.assembler.segments.pop(head_sentence_id, None) is not None:
# 与静音提交一致final 无文本时撤回该 id 的 partial不留下悬空等待。
await self.emit_state()
del self.segment_audio[:split_at]
self.segment_start_ms = head_end_ms
self.segment_id += 1
self.silence_ms = 0
self.voiced_ms = tail_voiced_ms
self.turn_reference = new_reference
self.window_start_bytes = 0
self.window_voiced_ms = 0.0
async def process_audio(self) -> None:
"""消费音频,以 VAD 静音结束作为切句主逻辑,并按窗口发送 partial。"""
last_partial_bytes = 0
partial_bytes = int(self.partial_interval_ms / 1000 * PARTIAL_BYTES_PER_SECOND)
while True:
item = await self.audio_queue.get()
if isinstance(item, EndOfStream):
break
chunk = self._strip_wav_header(item) if self.source == "file" else item
self.metrics.input_chunks += 1
if not self.windowed_partial:
# websocket_handler 已拒绝压缩文件;这里保留防御分支,避免未来
# 新客户端绕过入口时又悄悄退化成“整段上传后切片”。
raise RuntimeError("实时流式模式只接受 16kHz PCM16 音频")
if not chunk:
continue
self.vad_buffer.extend(chunk)
while len(self.vad_buffer) >= VAD_FRAME_BYTES:
frame = bytes(self.vad_buffer[:VAD_FRAME_BYTES])
del self.vad_buffer[:VAD_FRAME_BYTES]
frame_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
self.processed_audio_bytes += len(frame)
self.metrics.audio_bytes += len(frame)
voiced = self._is_voice_frame(frame)
if voiced and not self.in_speech:
self.in_speech = True
self.segment_start_ms = frame_start_ms - len(self.pre_roll) / PARTIAL_BYTES_PER_SECOND * 1000
self.segment_audio = bytearray(self.pre_roll)
self.pre_roll.clear()
last_partial_bytes = 0
self._reset_turn_embedding_state()
if self.in_speech:
self.segment_audio.extend(frame)
if voiced:
self.voiced_ms += VAD_FRAME_MS
self.window_voiced_ms += VAD_FRAME_MS
self.silence_ms = 0 if voiced else self.silence_ms + VAD_FRAME_MS
if len(self.segment_audio) - last_partial_bytes >= partial_bytes and self.silence_ms < self.silence_limit_ms:
text = await self._transcribe(partial=True)
if text:
await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms())
last_partial_bytes = len(self.segment_audio)
if (self.window_check_available
and len(self.segment_audio) - self.window_start_bytes >= self.turn_window_bytes
and self.window_voiced_ms >= SPEAKER_WINDOW_MIN_VOICED_MS):
if await self._check_speaker_window():
# 尾段是全新文本,下一帧就允许触发它的 partial。
last_partial_bytes = 0
if self.silence_ms >= self.silence_limit_ms or self._duration_ms() >= self.max_segment_sec * 1000:
await self._commit_segment("silence" if self.silence_ms >= self.silence_limit_ms else "max_duration")
last_partial_bytes = 0
else:
# 参考原 WebSocket 保留 200ms 前滚,减少首字低能量音素被裁掉。
self.pre_roll.extend(frame)
del self.pre_roll[:-6400]
if not self.wav_payload_started or (not self.input_stopped and self.wav_data_remaining not in (None, 0)):
raise ValueError("WAV 文件不完整,未收到全部音频数据")
if len(self.vad_buffer) % 2:
raise ValueError("PCM16 音频必须包含完整的双字节采样")
if self.windowed_partial and self.vad_buffer:
tail_ms = len(self.vad_buffer) / PARTIAL_BYTES_PER_SECOND * 1000
tail_voiced = self._is_voice_frame(bytes(self.vad_buffer))
if tail_voiced and not self.in_speech:
self.in_speech = True
self.segment_start_ms = self.processed_audio_bytes / PARTIAL_BYTES_PER_SECOND * 1000
self.processed_audio_bytes += len(self.vad_buffer)
self.metrics.audio_bytes += len(self.vad_buffer)
if self.in_speech:
self.segment_audio.extend(self.vad_buffer)
self.silence_ms = 0 if tail_voiced else self.silence_ms + tail_ms
self.voiced_ms += tail_ms if tail_voiced else 0
self.vad_buffer.clear()
if self.segment_audio:
if not self.windowed_partial:
self.in_speech = True
await self._commit_segment()
await self.emit({"type": "metrics", "metrics": self.metrics.snapshot()})
def deployment_model_name() -> str:
"""把部署脚本的 0.6b/1.7b 别名解析成 vLLM 对外发布的模型名。"""
public_name = os.getenv("VLLM_SERVED_MODEL_NAME")
if public_name:
return public_name
requested = os.getenv("QWEN3_ASR_MODEL", "default")
with (DEPLOY_ROOT / "model_manifest.json").open(encoding="utf-8") as source:
manifest = json.load(source)
if requested == "default":
return str(manifest["default_model"])
for model_id, config in manifest["models"].items():
if requested.lower() == str(config.get("alias", "")).lower():
return model_id
return requested
def parse_args() -> argparse.Namespace:
"""只解析服务选择参数;浏览器服务端口继续由代码内部变量统一维护。"""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-service-url", default=os.getenv("MODEL_SERVICE_URL", DEFAULT_MODEL_SERVICE_URL))
parser.add_argument("--model", default=deployment_model_name())
parser.add_argument("--no-browser", action="store_true")
return parser.parse_args()
async def index_handler(_: web.Request) -> web.FileResponse:
"""返回独立 Demo 测试页面,并避免入口页缓存旧的脚本版本号。"""
# 入口页必须每次重新校验,配合 app.js 的版本号变更,避免用户继续运行旧前端。
return web.FileResponse(
Path(__file__).parent / "static" / "index.html",
headers={"Cache-Control": "no-store"},
)
async def config_handler(request: web.Request) -> web.Response:
"""暴露服务启动时的默认配置,让页面自动填充 VLLM 地址和模型名。"""
config = request.app[MODEL_SERVICE_KEY].config
auxiliary = request.app.get(AUXILIARY_SERVICE_KEY)
return web.json_response({
"model_service_url": config.base_url, "model": config.model,
"speaker_service_url": getattr(getattr(auxiliary, "config", None), "base_url", None),
})
async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
"""处理一个浏览器会话,每个连接独立保存音频、句子和展示状态。"""
ws = web.WebSocketResponse(max_msg_size=64 * 1024 * 1024)
await ws.prepare(request)
default_model_service: VLLMTranscriptionService = request.app[MODEL_SERVICE_KEY]
model_service = default_model_service
owns_model_service = False
processing: asyncio.Task[None] | None = None
speaker_processing: asyncio.Task[None] | None = None
session: RealtimeSession | None = None
try:
first = await ws.receive()
if first.type != WSMsgType.TEXT:
await ws.send_json({"type": "error", "message": "first message must be JSON start"})
return ws
try:
start = json.loads(first.data)
except json.JSONDecodeError:
await ws.send_json({"type": "error", "message": "invalid start JSON"})
return ws
if not isinstance(start, dict) or start.get("type") != "start":
await ws.send_json({"type": "error", "message": "first message must have type=start"})
return ws
# 本项目用于验证实时流式链路,文件模式只接受可以按 PCM 帧连续处理的
# WAV/PCMMP3、M4A 等压缩容器只能在文件完整到达后解码,不纳入本次测试。
source = str(start.get("source") or "mic")
file_suffix = Path(str(start.get("file_name") or "")).suffix.lower()
if source == "file" and file_suffix not in {".pcm", ".wav"}:
await ws.send_json(
{
"type": "error",
"message": "实时流式测试的文件模式只支持 PCM 或 WAV请改用麦克风、PCM 或 WAV",
}
)
return ws
# 页面可以在不重启 WebSocket Demo 的情况下为当前会话切换 VLLM 地址;
# 未切换时继续复用默认服务,避免普通场景为每个连接重复创建 HTTP 会话。
try:
requested_url = validate_model_service_url(
str(start.get("model_service_url") or default_model_service.config.base_url)
)
except ValueError as exc:
await ws.send_json({"type": "error", "message": str(exc)})
return ws
requested_model = str(start.get("model") or default_model_service.config.model).strip()
if (
requested_url != default_model_service.config.base_url
or requested_model != default_model_service.config.model
):
model_service = VLLMTranscriptionService(
ModelServiceConfig(base_url=requested_url, model=requested_model)
)
await model_service.start()
owns_model_service = True
auxiliary_service = request.app.get(AUXILIARY_SERVICE_KEY)
session = RealtimeSession(ws, model_service, auxiliary_service, start)
auxiliary_config = getattr(auxiliary_service, "config", None)
speaker_health: dict[str, Any] | None = None
speaker_health_error: str | None = None
if session.speaker_enabled and auxiliary_service is not None:
# 健康检查只用于尽早暴露辅助服务问题;即使失败也不阻断 ASR
# 这样可以从同一页面继续观察 ASR 与说话人链路的差异。
health_check = getattr(auxiliary_service, "health", None)
if callable(health_check):
try:
speaker_health = await asyncio.wait_for(health_check(), timeout=5)
if speaker_health.get("speaker_embedding_ready", speaker_health.get("ready")) is False:
speaker_health_error = "辅助模型服务未就绪,请检查 /health 返回的 models 状态"
except Exception as exc:
speaker_health_error = f"说话人辅助服务不可用:{exc}"
await session.emit(
{
"type": "start",
"model_service_url": model_service.config.base_url,
"model": model_service.config.model,
"session_id": session.session_id,
"enable_native_partial_stream": session.enable_native_partial,
"native_partial_supported": model_service.native_partial_supported,
"partial_mode": "http_cumulative_window",
"speaker_diarization_enabled": session.speaker_enabled,
"speaker_service_url": getattr(auxiliary_config, "base_url", None),
"speaker_service_health": speaker_health,
"sentence_strategy": session.sentence_strategy,
"silence_limit_ms": session.silence_limit_ms,
"display_state_supported": True,
}
)
if speaker_health_error:
await session.warn_speaker(speaker_health_error)
processing = asyncio.create_task(session.process_audio())
if session.speaker_enabled:
# speaker worker 与音频处理并行运行;它只消费已经结束的 turn
# 因此不会阻塞下一帧音频进入队列或影响 ASR partial 输出。
speaker_processing = asyncio.create_task(session.process_speakers())
async def guarded(operation):
"""接收和队列背压同时监听 worker推理失败立即报错而非永远等 stop。"""
pending = asyncio.create_task(operation)
try:
workers = [task for task in (processing, speaker_processing) if task is not None]
done, _ = await asyncio.wait([pending, *workers], return_when=asyncio.FIRST_COMPLETED)
if pending in done:
return await pending
for worker in workers:
if worker in done:
await worker
raise RuntimeError("实时处理任务意外结束")
return await pending
finally:
if not pending.done():
pending.cancel()
await asyncio.gather(pending, return_exceptions=True)
input_finished = False
while not ws.closed:
message = await guarded(ws.receive())
if message.type == WSMsgType.BINARY:
await guarded(session.audio_queue.put(bytes(message.data)))
continue
if message.type == WSMsgType.TEXT:
try:
control = json.loads(message.data)
except json.JSONDecodeError:
continue
if not isinstance(control, dict):
continue
if control.get("type") in {"eof", "stop"}:
input_finished = True
session.input_stopped = control.get("type") == "stop"
await session.emit({"type": "draining", "message": "正在完成转写和说话人识别"})
await guarded(session.audio_queue.put(EOF))
break
if control.get("type") == "abort":
processing.cancel()
if speaker_processing is not None:
speaker_processing.cancel()
await asyncio.gather(
processing,
*(task for task in [speaker_processing] if task is not None),
return_exceptions=True,
)
return ws
if message.type in {WSMsgType.ERROR, WSMsgType.CLOSE, WSMsgType.CLOSED}:
processing.cancel()
if speaker_processing is not None:
speaker_processing.cancel()
await asyncio.gather(
processing,
*(task for task in [speaker_processing] if task is not None),
return_exceptions=True,
)
return ws
if not input_finished:
# 浏览器或网络断开后,音频生产者已经不存在,不能让处理任务继续等待
# 永远不会到来的 EOF因此这里主动取消任务并回收异常结果。
processing.cancel()
if speaker_processing is not None:
speaker_processing.cancel()
await asyncio.gather(
processing,
*(task for task in [speaker_processing] if task is not None),
return_exceptions=True,
)
return ws
try:
await processing
if speaker_processing is not None:
await session.speaker_queue.put(EOF)
await speaker_processing
await session.emit_state()
await session.emit({
"type": "end", "metrics": session.metrics.snapshot(),
"sentences": session.assembler.raw_snapshot(),
"display_blocks": session.assembler.display_blocks(session.merge_adjacent),
})
except asyncio.CancelledError:
raise
except Exception as exc:
await session.emit({"type": "error", "message": str(exc)})
except Exception as exc:
LOGGER.exception("WebSocket session failed")
if not ws.closed:
await ws.send_json({"type": "error", "message": str(exc)})
finally:
for task in (processing, speaker_processing):
if task is not None and not task.done():
task.cancel()
pending_tasks = [task for task in (processing, speaker_processing) if task is not None]
if pending_tasks:
await asyncio.gather(*pending_tasks, return_exceptions=True)
# stop、abort、断线和推理异常均释放会话清理失败不覆盖最终识别结果。
if session is not None:
reset = getattr(session.auxiliary_service, "reset_speaker_session", None)
if reset is not None:
try:
await asyncio.wait_for(reset(session.session_id), timeout=5)
except Exception:
LOGGER.warning("speaker session cleanup failed: %s", session.session_id, exc_info=True)
if owns_model_service:
await model_service.close()
if not ws.closed:
await ws.close()
return ws
async def start_app(model_service_url: str, model: str) -> web.Application:
"""创建 HTTP/WebSocket 应用,并挂载可复用的 VLLM 适配器。"""
app = web.Application()
app[MODEL_SERVICE_KEY] = VLLMTranscriptionService(ModelServiceConfig(base_url=model_service_url, model=model))
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
AuxiliaryServiceConfig(base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"))
)
async def lifecycle(application: web.Application):
await application[MODEL_SERVICE_KEY].start()
await application[AUXILIARY_SERVICE_KEY].start()
yield
await application[AUXILIARY_SERVICE_KEY].close()
await application[MODEL_SERVICE_KEY].close()
app.cleanup_ctx.append(lifecycle)
app.router.add_get("/", index_handler)
app.router.add_get("/api/config", config_handler)
app.router.add_static("/static/", Path(__file__).parent / "static")
app.router.add_get("/ws", websocket_handler)
# 参考腾讯 Demo 将 static 目录挂载到根路径;页面中的 style.css 和 app.js
# 使用相对地址,必须同时提供根路径静态资源路由,否则浏览器会显示无样式页面。
app.router.add_static("/", Path(__file__).parent / "static", show_index=False)
return app
def main() -> None:
"""启动本地测试页面和 WebSocket 服务。"""
args = parse_args()
logging.basicConfig(level=logging.INFO)
if not args.no_browser:
webbrowser.open(f"http://{WEB_DISPLAY_HOST}:{WEB_PORT}/")
print(f"WebSocket demo: http://{WEB_DISPLAY_HOST}:{WEB_PORT}/", flush=True)
print(f"VLLM service: {args.model_service_url} ({args.model})", flush=True)
web.run_app(start_app(args.model_service_url, args.model), host=WEB_HOST, port=WEB_PORT)
if __name__ == "__main__":
main()