"""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 backend/requirements-funasr.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]