ASR-demo/backend/realtime_websocket/funasr_engine.py

394 lines
15 KiB
Python

"""FunASR realtime engine migrated into the demo project.
This module keeps the browser-facing project independent from the original
Qwen/VLLM service. It follows FunASR's streaming cache lifecycle:
one cache per ASR session and is_final=True only for the last chunk.
"""
from __future__ import annotations
import asyncio
import logging
import os
from dataclasses import dataclass
from typing import Any
LOGGER = logging.getLogger(__name__)
@dataclass(frozen=True)
class FunASRServiceConfig:
"""Configuration for the in-process FunASR models."""
model: str = "paraformer-zh-streaming"
vad_model: str = "fsmn-vad"
device: str = "cuda:0"
vad_device: str = "cpu"
sample_rate: int = 16000
vad_chunk_ms: int = 200
chunk_size: tuple[int, int, int] = (0, 10, 5)
encoder_chunk_look_back: int = 4
decoder_chunk_look_back: int = 1
max_segment_sec: float = 30.0
@classmethod
def from_env(cls) -> "FunASRServiceConfig":
"""Read model selection from environment without changing WS fields."""
chunk_text = os.getenv("FUNASR_CHUNK_SIZE", "0,10,5")
try:
values = tuple(int(value.strip()) for value in chunk_text.split(","))
chunk_size = values if len(values) == 3 else cls.chunk_size
except ValueError:
chunk_size = cls.chunk_size
return cls(
model=os.getenv("FUNASR_ASR_MODEL", cls.model),
vad_model=os.getenv("FUNASR_VAD_MODEL", cls.vad_model),
device=os.getenv("FUNASR_DEVICE", cls.device),
vad_device=os.getenv("FUNASR_VAD_DEVICE", cls.vad_device),
vad_chunk_ms=max(50, int(os.getenv("FUNASR_VAD_CHUNK_MS", str(cls.vad_chunk_ms)))),
chunk_size=chunk_size,
encoder_chunk_look_back=max(
0, int(os.getenv("FUNASR_ENCODER_LOOK_BACK", str(cls.encoder_chunk_look_back)))
),
decoder_chunk_look_back=max(
0, int(os.getenv("FUNASR_DECODER_LOOK_BACK", str(cls.decoder_chunk_look_back)))
),
max_segment_sec=max(
2.0, float(os.getenv("FUNASR_MAX_SEGMENT_SEC", str(cls.max_segment_sec)))
),
)
@dataclass(frozen=True)
class FunASRSegment:
"""A partial or final event returned by one browser session."""
text: str
start_time_ms: float
end_time_ms: float
audio: bytes
voiced_ms: float
is_final: bool
sentence_id: int
reason: str | None = None
def _result_text(result: Any) -> str:
"""Extract text from the list/dict result shapes used by FunASR."""
if isinstance(result, list):
return _result_text(result[0]) if result else ""
if isinstance(result, dict):
value = result.get("text")
if value is None:
value = result.get("value")
return str(value or "").strip()
return str(result or "").strip()
def _vad_events(result: Any) -> list[tuple[float, float]]:
"""Normalize FunASR streaming VAD output to start/end milliseconds."""
if isinstance(result, list):
result = result[0] if result else {}
if isinstance(result, dict):
result = result.get("value") or result.get("segments") or []
if not isinstance(result, list):
return []
events: list[tuple[float, float]] = []
for item in result:
if not isinstance(item, (list, tuple)) or len(item) < 2:
continue
try:
events.append((float(item[0]), float(item[1])))
except (TypeError, ValueError):
continue
return events
class FunASRModelService:
"""Load FunASR once and create isolated streaming state per browser."""
native_partial_supported = True
def __init__(self, config: FunASRServiceConfig | None = None) -> None:
self.config = config or FunASRServiceConfig.from_env()
self.asr_model: Any | None = None
self.vad_model: Any | None = None
# Model objects are shared; each session owns its own cache. Serializing
# calls makes the first validation version predictable on one GPU.
self.inference_lock = asyncio.Lock()
async def start(self) -> None:
"""Load streaming ASR and VAD outside the event loop."""
try:
from funasr import AutoModel
except ImportError as exc: # pragma: no cover - deployment-only branch
raise RuntimeError(
"FunASR is not installed; run pip install -r requirements.txt"
) from exc
def load_models() -> tuple[Any, Any]:
common = {"disable_pbar": True, "disable_log": True}
asr = AutoModel(model=self.config.model, device=self.config.device, **common)
vad = AutoModel(model=self.config.vad_model, device=self.config.vad_device, **common)
return asr, vad
self.asr_model, self.vad_model = await asyncio.to_thread(load_models)
LOGGER.info(
"FunASR ready: model=%s vad=%s device=%s vad_device=%s",
self.config.model,
self.config.vad_model,
self.config.device,
self.config.vad_device,
)
async def close(self) -> None:
"""Release references so CUDA memory can be reclaimed on shutdown."""
self.asr_model = None
self.vad_model = None
def create_session(self) -> "FunASRRealtimeSession":
"""Create a session with isolated VAD and ASR caches."""
if self.asr_model is None or self.vad_model is None:
raise RuntimeError("FunASR model service is not started")
return FunASRRealtimeSession(self)
async def generate_asr(self, audio: Any, status: dict[str, Any]) -> str:
"""Run one blocking ASR chunk while preserving its mutable cache."""
if self.asr_model is None:
raise RuntimeError("FunASR ASR model is not loaded")
def generate() -> str:
return _result_text(self.asr_model.generate(input=audio, **status))
async with self.inference_lock:
return await asyncio.to_thread(generate)
async def generate_vad(
self,
audio: Any,
status: dict[str, Any],
chunk_ms: int,
) -> list[tuple[float, float]]:
"""Run one streaming VAD chunk and normalize endpoint events."""
if self.vad_model is None:
raise RuntimeError("FunASR VAD model is not loaded")
def generate() -> list[tuple[float, float]]:
result = self.vad_model.generate(input=audio, chunk_size=chunk_ms, **status)
return _vad_events(result)
async with self.inference_lock:
return await asyncio.to_thread(generate)
class _StreamingASRTurn:
"""One utterance using FunASR's ordered chunk/cache lifecycle."""
def __init__(self, service: FunASRModelService) -> None:
self.service = service
self.cache: dict[str, Any] = {}
self.pending = bytearray()
self.cumulative_text = ""
self.last_chunk_text = ""
self.pending_outputs: list[str] = []
self.chunk_samples = max(1, service.config.chunk_size[1] * 960)
@staticmethod
def _merge_chunk_text(current: str, chunk: str, previous_chunk: str) -> str:
"""Merge chunk text while tolerating wrappers returning cumulative text."""
if not chunk or chunk == previous_chunk:
return current
if current and chunk.startswith(current):
return chunk
return current + chunk
async def append(self, pcm_bytes: bytes) -> list[str]:
"""Decode complete chunks but hold one chunk for the final flush."""
if not pcm_bytes:
return []
self.pending.extend(pcm_bytes)
chunk_bytes = self.chunk_samples * 2
# Hold one full chunk so the real last chunk receives is_final=True.
while len(self.pending) >= chunk_bytes * 2:
chunk = bytes(self.pending[:chunk_bytes])
del self.pending[:chunk_bytes]
text = await self.service.generate_asr(
chunk,
{
"cache": self.cache,
"is_final": False,
"chunk_size": self.service.config.chunk_size,
"encoder_chunk_look_back": self.service.config.encoder_chunk_look_back,
"decoder_chunk_look_back": self.service.config.decoder_chunk_look_back,
"batch_size": 1,
},
)
self.cumulative_text = self._merge_chunk_text(
self.cumulative_text, text, self.last_chunk_text
)
self.last_chunk_text = text
if text:
self.pending_outputs.append(self.cumulative_text)
outputs = self.pending_outputs
self.pending_outputs = []
return outputs
async def finish(self) -> str:
"""Flush the last buffered chunk and return cumulative text."""
if self.pending:
chunk = bytes(self.pending)
self.pending.clear()
text = await self.service.generate_asr(
chunk,
{
"cache": self.cache,
"is_final": True,
"chunk_size": self.service.config.chunk_size,
"encoder_chunk_look_back": self.service.config.encoder_chunk_look_back,
"decoder_chunk_look_back": self.service.config.decoder_chunk_look_back,
"batch_size": 1,
},
)
self.cumulative_text = self._merge_chunk_text(
self.cumulative_text, text, self.last_chunk_text
)
self.last_chunk_text = text
return self.cumulative_text.strip()
class FunASRRealtimeSession:
"""FunASR VAD + streaming ASR session used by one browser WebSocket."""
def __init__(self, service: FunASRModelService) -> None:
self.service = service
self.sample_rate = service.config.sample_rate
self.vad_chunk_bytes = service.config.vad_chunk_ms * self.sample_rate * 2 // 1000
self.vad_buffer = bytearray()
self.vad_cache: dict[str, Any] = {}
self.pre_roll = bytearray()
self.pre_roll_max_bytes = 300 * self.sample_rate * 2 // 1000
self.total_samples = 0
self.speech_started = False
self.segment_id = 0
self.segment_start_ms = 0.0
self.segment_audio = bytearray()
self.segment_voiced_ms = 0.0
self.asr_turn: _StreamingASRTurn | None = None
@staticmethod
def _to_float32(pcm_bytes: bytes) -> Any:
"""Convert browser PCM16 to the float waveform expected by FunASR."""
import numpy as np
return np.frombuffer(pcm_bytes, dtype=np.int16).astype(np.float32) / 32768.0
def _start_segment(self, start_ms: float) -> None:
"""Open a turn and retain a short pre-roll for initial phonemes."""
self.speech_started = True
self.segment_start_ms = max(
0.0,
start_ms - len(self.pre_roll) / (self.sample_rate * 2) * 1000,
)
self.segment_audio = bytearray(self.pre_roll)
self.segment_voiced_ms = 0.0
self.asr_turn = _StreamingASRTurn(self.service)
self.pre_roll.clear()
async def _feed_vad_chunk(self, chunk: bytes) -> list[FunASRSegment]:
"""Feed VAD, then stream this audio into the active ASR turn."""
self.total_samples += len(chunk) // 2
events = await self.service.generate_vad(
self._to_float32(chunk),
{"cache": self.vad_cache, "is_final": False},
self.service.config.vad_chunk_ms,
)
starts = [start for start, _ in events if start >= 0]
ends = [end for _, end in events if end >= 0]
results: list[FunASRSegment] = []
if not self.speech_started and starts:
self._start_segment(starts[0])
if self.speech_started:
self.segment_audio.extend(chunk)
self.segment_voiced_ms += len(chunk) / (self.sample_rate * 2) * 1000
if self.asr_turn is not None:
for text in await self.asr_turn.append(chunk):
results.append(self._partial(text))
if self._segment_duration_ms() >= self.service.config.max_segment_sec * 1000:
results.append(await self._finish_segment("max_duration"))
else:
self.pre_roll.extend(chunk)
del self.pre_roll[:-self.pre_roll_max_bytes]
if ends and self.speech_started:
results.append(await self._finish_segment("vad_end"))
return [event for event in results if event.text or event.is_final]
def _segment_duration_ms(self) -> float:
return len(self.segment_audio) / (self.sample_rate * 2) * 1000
def _partial(self, text: str) -> FunASRSegment:
"""Create a partial event using the cumulative FunASR text."""
return FunASRSegment(
text=text,
start_time_ms=self.segment_start_ms,
end_time_ms=self.segment_start_ms + self._segment_duration_ms(),
audio=b"",
voiced_ms=self.segment_voiced_ms,
is_final=False,
sentence_id=self.segment_id,
)
async def _finish_segment(self, reason: str) -> FunASRSegment:
"""Flush the ASR cache, then release completed segment audio."""
text = await self.asr_turn.finish() if self.asr_turn is not None else ""
result = FunASRSegment(
text=text,
start_time_ms=self.segment_start_ms,
end_time_ms=self.segment_start_ms + self._segment_duration_ms(),
audio=bytes(self.segment_audio),
voiced_ms=self.segment_voiced_ms,
is_final=True,
sentence_id=self.segment_id,
reason=reason,
)
self.segment_id += 1
self.segment_audio.clear()
self.asr_turn = None
self.speech_started = False
self.segment_voiced_ms = 0.0
return result
async def feed(self, pcm_bytes: bytes) -> list[FunASRSegment]:
"""Consume PCM16 and return FunASR partial/final events."""
if len(pcm_bytes) % 2:
raise ValueError("PCM16 音频必须包含完整的双字节采样")
self.vad_buffer.extend(pcm_bytes)
results: list[FunASRSegment] = []
while len(self.vad_buffer) >= self.vad_chunk_bytes:
chunk = bytes(self.vad_buffer[:self.vad_chunk_bytes])
del self.vad_buffer[:self.vad_chunk_bytes]
results.extend(await self._feed_vad_chunk(chunk))
return results
async def finish(self) -> list[FunASRSegment]:
"""Flush buffered audio and discard all stream caches."""
results: list[FunASRSegment] = []
if self.vad_buffer:
chunk = bytes(self.vad_buffer)
self.vad_buffer.clear()
if not self.speech_started and chunk:
self._start_segment(self.total_samples / self.sample_rate * 1000)
self.total_samples += len(chunk) // 2
if self.speech_started:
self.segment_audio.extend(chunk)
self.segment_voiced_ms += len(chunk) / (self.sample_rate * 2) * 1000
if self.asr_turn is not None:
for text in await self.asr_turn.append(chunk):
results.append(self._partial(text))
if self.speech_started:
results.append(await self._finish_segment("eof"))
self.vad_cache = {}
self.pre_roll.clear()
return [event for event in results if event.text or event.is_final]