912 lines
44 KiB
Python
912 lines
44 KiB
Python
"""面向浏览器的独立 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/PCM;MP3、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()
|