修改ASR采用流式返回模式

main
Bifang 2026-09-22 16:50:00 +08:00
parent aeaa72fc63
commit 525d8060f5
7 changed files with 305 additions and 13 deletions

View File

@ -84,4 +84,4 @@ python -m unittest discover -s tests -v
node --test tests/test_frontend.cjs
```
这些测试覆盖状态机、模拟 HTTP/WebSocket、延迟更新、前端脚本和有效向量校验,不执行模型下载或 GPU 推理。当前 VLLM 适配器仍是 HTTP 累积窗口 partial,`native_partial_supported=false`;本次没有把 HTTP 接口包装成原生增量模型状态。单个内部片段中无停顿的多人换话或重叠讲话仍需真实模型和更细粒度切段验证。
这些测试覆盖状态机、模拟 HTTP/WebSocket、延迟更新、前端脚本和有效向量校验,不执行模型下载或 GPU 推理。当前 VLLM 适配器默认使用 `/v1/realtime` 原生流式,旧同步 HTTP 保留为兼容回退。单个内部片段中无停顿的多人换话或重叠讲话仍需真实模型和更细粒度切段验证。

View File

@ -15,11 +15,11 @@
本次修复、诊断状态和部署验收步骤见 [FIXES.md](FIXES.md)。ASR 继续使用已部署的独立 vLLM;已有 vLLM 服务时无需重复启动或下载模型。
当前 VLLM 端点提供的是同步 OpenAI 音频转写接口,没有暴露原生
`create_stream/feed_stream/finish_stream`。因此本项目仍然是真实 WebSocket
音频流:麦克风 PCM 到达后立即进入 VAD,按窗口调用 VLLM 生成 partial;它不会
等整段音频结束。`native_partial_supported=false` 只表示模型 HTTP 接口本身
不是原生 ASR stream,不伪造不存在的能力。
当前默认使用 vLLM `/v1/realtime` WebSocket:PCM 音频帧只发送一次,模型
持续返回 `transcription.delta`,VAD 切段时发送 `input_audio_buffer.commit`。
因此 partial 不再重复上传不断增长的整段音频,适合长时间会议运行。
启动脚本默认启用 `Qwen3ASRRealtimeGeneration`;如果部署端点不支持 realtime,
本项目会自动回退到原同步 HTTP 接口,外部 WebSocket 参数不变。
## 启动

View File

@ -2,13 +2,17 @@
from __future__ import annotations
import asyncio
import base64
import io
import json
import wave
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlsplit, urlunsplit
from aiohttp import ClientSession, ClientTimeout, FormData
from aiohttp import ClientSession, ClientTimeout, FormData, WSMsgType
@dataclass(frozen=True)
@ -19,6 +23,7 @@ class ModelServiceConfig:
model: str = "Qwen/Qwen3-ASR-0.6B"
api_key: str = "EMPTY"
timeout_seconds: float = 45.0
realtime_enabled: bool = True
def pcm16_to_wav(pcm_bytes: bytes, sample_rate: int = 16000) -> bytes:
@ -66,11 +71,113 @@ def prepare_audio_request(
}.get(suffix, "application/octet-stream")
return audio_bytes, Path(file_name).name or "audio.bin", content_type
def realtime_ws_url(base_url: str) -> str:
"""Convert the configured OpenAI-compatible base URL to vLLM's realtime URL."""
parsed = urlsplit(base_url.rstrip("/"))
if parsed.scheme not in {"http", "https"} or not parsed.netloc:
raise ValueError(f"invalid VLLM base URL: {base_url}")
scheme = "wss" if parsed.scheme == "https" else "ws"
path = parsed.path.rstrip("/") + "/realtime"
return urlunsplit((scheme, parsed.netloc, path, "", ""))
class VLLMRealtimeStream:
"""One vLLM realtime stream, isolated from the shared HTTP client session."""
def __init__(self, websocket: Any, timeout_seconds: float) -> None:
self._websocket = websocket
self._timeout_seconds = timeout_seconds
self._latest_text = ""
self._error: Exception | None = None
self._done = asyncio.Event()
self._reader = asyncio.create_task(self._read_messages())
async def _read_messages(self) -> None:
"""Collect model deltas continuously so audio ingestion never waits for a snapshot."""
try:
async for message in self._websocket:
if message.type == WSMsgType.TEXT:
try:
payload = json.loads(message.data)
except (TypeError, ValueError):
continue
event_type = payload.get("type")
if event_type == "transcription.delta":
delta = str(payload.get("delta") or "")
if delta:
self._latest_text += delta
elif payload.get("text") is not None:
self._latest_text = str(payload["text"])
elif event_type == "transcription.done":
self._latest_text = str(
payload.get("text") or payload.get("transcript") or self._latest_text
).strip()
self._done.set()
elif event_type == "error":
detail = payload.get("error") or payload.get("message") or "unknown realtime error"
self._error = RuntimeError(str(detail))
self._done.set()
elif message.type == WSMsgType.ERROR:
self._error = self._websocket.exception() or RuntimeError("VLLM realtime WebSocket failed")
self._done.set()
return
elif message.type in {WSMsgType.CLOSE, WSMsgType.CLOSED, WSMsgType.CLOSING}:
if not self._done.is_set():
self._error = RuntimeError("VLLM realtime WebSocket closed before transcription.done")
self._done.set()
return
except asyncio.CancelledError:
raise
except Exception as exc:
self._error = exc
self._done.set()
def latest_text(self) -> str:
"""Return the newest model text already received by the reader task."""
return self._latest_text.strip()
async def append_audio(self, pcm_bytes: bytes) -> None:
"""Push one raw 16 kHz mono PCM16 block without re-uploading old audio."""
if not pcm_bytes:
return
if self._error is not None:
raise self._error
await self._websocket.send_json(
{
"type": "input_audio_buffer.append",
"audio": base64.b64encode(pcm_bytes).decode("ascii"),
}
)
async def finish(self) -> str:
"""Commit the current model turn and return the final realtime transcription."""
if self._error is not None:
raise self._error
await self._websocket.send_json({"type": "input_audio_buffer.commit", "final": True})
try:
await asyncio.wait_for(self._done.wait(), timeout=self._timeout_seconds)
except asyncio.TimeoutError as exc:
raise TimeoutError("VLLM realtime transcription timed out") from exc
if self._error is not None:
raise self._error
return self._latest_text.strip()
async def close(self) -> None:
"""Stop the reader and release the model-side WebSocket."""
if not self._reader.done():
self._reader.cancel()
await asyncio.gather(self._reader, return_exceptions=True)
if not self._websocket.closed:
await self._websocket.close()
class VLLMTranscriptionService:
"""只调用独立项目提供的 VLLM HTTP 接口,不导入原项目应用代码。"""
native_partial_supported = False
@property
def native_partial_supported(self) -> bool:
"""Report whether this adapter is configured to use vLLM realtime."""
return self.config.realtime_enabled
def __init__(self, config: ModelServiceConfig) -> None:
self.config = config
@ -86,6 +193,35 @@ class VLLMTranscriptionService:
await self._session.close()
self._session = None
async def open_realtime_stream(self) -> VLLMRealtimeStream:
"""Open a model-native stream; the caller owns and closes the returned turn."""
if not self.config.realtime_enabled:
raise RuntimeError("VLLM realtime streaming is disabled")
if self._session is None:
raise RuntimeError("model service is not started")
endpoint = realtime_ws_url(self.config.base_url)
headers = {"Authorization": f"Bearer {self.config.api_key}"}
connect_timeout = min(5.0, self.config.timeout_seconds)
websocket = await self._session.ws_connect(
endpoint,
headers=headers,
timeout=connect_timeout,
heartbeat=20,
)
try:
created = await asyncio.wait_for(websocket.receive(), timeout=connect_timeout)
if created.type == WSMsgType.TEXT:
payload = json.loads(created.data)
if payload.get("type") == "error":
raise RuntimeError(
str(payload.get("error") or payload.get("message") or "VLLM realtime error")
)
await websocket.send_json({"type": "session.update", "model": self.config.model})
return VLLMRealtimeStream(websocket, self.config.timeout_seconds)
except Exception:
await websocket.close()
raise
async def transcribe(
self,
audio_bytes: bytes,

View File

@ -136,6 +136,9 @@ class RealtimeSession:
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.native_stream: Any | None = None
self.native_stream_disabled = False
self.native_sent_bytes = 0
self.segment_id = 0
self.segment_audio = bytearray()
self.segment_start_ms = 0.0
@ -220,6 +223,70 @@ class RealtimeSession:
partial=partial,
)
async def _start_native_stream(self) -> None:
"""Start one model stream for the current VAD turn; HTTP remains the safe fallback."""
if not self.enable_native_partial or self.native_stream_disabled or self.native_stream is not None:
return
opener = getattr(self.model_service, "open_realtime_stream", None)
if not callable(opener):
self.native_stream_disabled = True
return
try:
self.native_stream = await opener()
self.native_sent_bytes = 0
except Exception as exc:
self.native_stream_disabled = True
LOGGER.warning("native ASR stream unavailable; falling back to HTTP: %s", exc)
async def _feed_native_audio(self) -> None:
"""Send only new PCM bytes so a long turn is never re-uploaded as a growing snapshot."""
if self.native_stream is None:
return
pending = bytes(self.segment_audio[self.native_sent_bytes:])
if not pending:
return
try:
await self.native_stream.append_audio(pending)
self.native_sent_bytes = len(self.segment_audio)
except Exception as exc:
await self._close_native_stream()
self.native_stream_disabled = True
LOGGER.warning("native ASR stream failed; falling back to HTTP: %s", exc)
async def _partial_text(self) -> str | None:
"""Read the latest native delta; use the old HTTP path only when native streaming is unavailable."""
if self.native_stream is not None:
return self.native_stream.latest_text()
return await self._transcribe(partial=True)
async def _finish_transcription(self) -> str | None:
"""Commit the native turn once, then fall back to one final HTTP request on failure."""
stream = self.native_stream
self.native_stream = None
self.native_sent_bytes = 0
if stream is None:
return await self._transcribe(partial=False)
try:
return await stream.finish()
except Exception as exc:
self.native_stream_disabled = True
LOGGER.warning("native ASR finalization failed; falling back to HTTP: %s", exc)
return await self._transcribe(partial=False)
finally:
await stream.close()
async def _close_native_stream(self) -> None:
"""Release an unfinished native turn when the browser disconnects or aborts."""
stream = self.native_stream
self.native_stream = None
self.native_sent_bytes = 0
if stream is not None:
await stream.close()
async def close(self) -> None:
"""Release the model stream owned by this browser session."""
await self._close_native_stream()
async def _emit_transcription(self, text: str, sentence_type: int, end_ms: float, commit_reason: str | None = None) -> None:
"""写入或更新一条句子,确保中间结果和最终结果不会在前端产生重复行。"""
if not text:
@ -373,7 +440,7 @@ class RealtimeSession:
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)
text = await self._finish_transcription()
if text:
await self._emit_transcription(text, 1, final_end_ms, reason)
if self.speaker_enabled:
@ -426,13 +493,16 @@ class RealtimeSession:
self.segment_audio = bytearray(self.pre_roll)
self.pre_roll.clear()
last_partial_bytes = 0
await self._start_native_stream()
await self._feed_native_audio()
if self.in_speech:
self.segment_audio.extend(frame)
await self._feed_native_audio()
if voiced:
self.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)
text = await self._partial_text()
if text:
await self._emit_transcription(text, 0, self.segment_start_ms + self._duration_ms())
last_partial_bytes = len(self.segment_audio)
@ -609,7 +679,7 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
"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",
"partial_mode": "vllm_realtime_websocket" if session.enable_native_partial and model_service.native_partial_supported else "http_cumulative_window",
"speaker_diarization_enabled": session.speaker_enabled,
"speaker_service_url": getattr(auxiliary_config, "base_url", None),
"speaker_service_health": speaker_health,
@ -725,6 +795,7 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
# stop、abort、断线和推理异常均释放会话,清理失败不覆盖最终识别结果。
if session is not None:
reset = getattr(session.auxiliary_service, "reset_speaker_session", None)
await session.close()
if reset is not None:
try:
await asyncio.wait_for(reset(session.session_id), timeout=5)
@ -740,7 +811,13 @@ async def websocket_handler(request: web.Request) -> web.WebSocketResponse:
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[MODEL_SERVICE_KEY] = VLLMTranscriptionService(
ModelServiceConfig(
base_url=model_service_url,
model=model,
realtime_enabled=os.getenv("VLLM_REALTIME", "true").lower() not in {"0", "false", "no", "off"},
)
)
app[AUXILIARY_SERVICE_KEY] = AuxiliaryModelService(
AuxiliaryServiceConfig(base_url=os.getenv("AUXILIARY_SERVICE_URL", "http://127.0.0.1:8010"))
)

View File

@ -2,14 +2,55 @@
from __future__ import annotations
import asyncio
import base64
import json
from types import SimpleNamespace
import unittest
import wave
from io import BytesIO
from model_service import pcm16_to_wav, prepare_audio_request, wav_to_pcm16
from model_service import (
pcm16_to_wav,
prepare_audio_request,
realtime_ws_url,
wav_to_pcm16,
VLLMRealtimeStream,
)
from auxiliary_service import AuxiliaryModelService, AuxiliaryServiceConfig
from aiohttp import WSMsgType
class FakeRealtimeWebSocket:
def __init__(self) -> None:
self.messages: asyncio.Queue[object] = asyncio.Queue()
self.sent: list[dict[str, object]] = []
self.closed = False
def __aiter__(self):
return self._messages()
async def _messages(self):
while not self.closed:
yield await self.messages.get()
async def send_json(self, payload: dict[str, object]) -> None:
self.sent.append(payload)
if payload.get("type") == "input_audio_buffer.commit" and payload.get("final"):
await self.messages.put(
SimpleNamespace(
type=WSMsgType.TEXT,
data=json.dumps({"type": "transcription.done", "text": "final text"}),
)
)
def exception(self):
return None
async def close(self) -> None:
self.closed = True
class ModelServiceTests(unittest.TestCase):
def test_pcm_is_wrapped_as_16k_mono_wav(self) -> None:
wav_bytes = pcm16_to_wav(b"\x00\x00" * 160)
@ -23,6 +64,34 @@ class ModelServiceTests(unittest.TestCase):
prepared = prepare_audio_request(b"complete", "file", "sample.mp3", partial=False)
self.assertEqual(prepared[1], "sample.mp3")
def test_realtime_url_uses_the_openai_v1_path(self) -> None:
self.assertEqual(
realtime_ws_url("https://asr.example/v1/"),
"wss://asr.example/v1/realtime",
)
def test_realtime_stream_sends_incremental_audio_and_finishes(self) -> None:
async def exercise() -> None:
websocket = FakeRealtimeWebSocket()
stream = VLLMRealtimeStream(websocket, timeout_seconds=1)
await websocket.messages.put(
SimpleNamespace(
type=WSMsgType.TEXT,
data=json.dumps({"type": "transcription.delta", "delta": "partial"}),
)
)
await stream.append_audio(b"\x01\x02")
final_text = await stream.finish()
self.assertEqual(final_text, "final text")
self.assertEqual(
base64.b64decode(str(websocket.sent[0]["audio"])),
b"\x01\x02",
)
self.assertEqual(websocket.sent[-1]["type"], "input_audio_buffer.commit")
await stream.close()
asyncio.run(exercise())
def test_auxiliary_config_is_independent_from_vllm(self) -> None:
service = AuxiliaryModelService(AuxiliaryServiceConfig())
self.assertEqual(service.config.base_url, "http://127.0.0.1:8010")

View File

@ -110,6 +110,12 @@ def add_arguments(parser: argparse.ArgumentParser) -> None:
action=argparse.BooleanOptionalAction,
default=os.getenv("VLLM_ENFORCE_EAGER", "true").lower() == "true",
)
parser.add_argument(
"--realtime",
action=argparse.BooleanOptionalAction,
default=os.getenv("VLLM_REALTIME", "true").lower() == "true",
help="Enable the vLLM realtime WebSocket architecture",
)
def build_server_command(args: argparse.Namespace, model_id: str, model_path: Path) -> list[str]:
@ -141,6 +147,8 @@ def build_server_command(args: argparse.Namespace, model_id: str, model_path: Pa
"--tensor-parallel-size",
str(args.tensor_parallel_size),
]
if args.realtime:
command.extend(["--hf-overrides", '{"architectures":["Qwen3ASRRealtimeGeneration"]}'])
if args.enforce_eager:
command.append("--enforce-eager")
return command

View File

@ -37,6 +37,8 @@ class ServeConfigTests(unittest.TestCase):
self.assertEqual(command[0:3], ["/opt/asr-gb10/bin/vllm", "serve", str(Path("/models/Qwen3-ASR-0.6B"))])
self.assertIn("--enforce-eager", command)
self.assertIn("--hf-overrides", command)
self.assertTrue(any("Qwen3ASRRealtimeGeneration" in item for item in command))
self.assertIn("9950", command)
if __name__ == "__main__":