修改ASR采用流式返回模式
parent
aeaa72fc63
commit
525d8060f5
|
|
@ -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 保留为兼容回退。单个内部片段中无停顿的多人换话或重叠讲话仍需真实模型和更细粒度切段验证。
|
||||
|
|
|
|||
|
|
@ -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 参数不变。
|
||||
|
||||
## 启动
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
Loading…
Reference in New Issue