"""FunASR 流式引擎适配器。 每个会话单独维护流式缓存,并且只将最后一个分块标记为结束块。 """ 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: """进程内 FunASR 模型的配置。""" 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": """从环境变量读取模型选择,不改变 WebSocket 字段。""" 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: """浏览器会话返回的中间结果或最终结果事件。""" 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: """从 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]]: """将 FunASR 流式 VAD 输出统一转换为毫秒起止时间。""" 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: """只加载一次 FunASR,并为每个浏览器会话创建独立的流状态。""" 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 # 模型对象由会话共享,每个会话单独持有缓存。串行处理调用, # 让首个验证版本在单卡环境中的行为更可预测。 self.inference_lock = asyncio.Lock() async def start(self) -> None: """在事件循环之外加载流式 ASR 和 VAD。""" 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: """释放对象引用,以便关闭时回收 CUDA 显存。""" self.asr_model = None self.vad_model = None def create_session(self) -> "FunASRRealtimeSession": """创建具有独立 VAD 和 ASR 缓存的会话。""" 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: """运行一个阻塞式 ASR 分块,同时保留其可变缓存。""" 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]]: """运行一个流式 VAD 分块,并统一端点事件格式。""" 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: """按照 FunASR 的分块与缓存顺序处理一段话语。""" 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: """合并分块文本,并兼容返回累计文本的封装实现。""" 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]: """解码完整分块,同时暂存一个分块用于最终刷新。""" if not pcm_bytes: return [] self.pending.extend(pcm_bytes) chunk_bytes = self.chunk_samples * 2 # 暂存一个完整分块,确保真正的最后一块收到 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: """刷新最后一个缓冲分块并返回累计文本。""" 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: """供单个浏览器 WebSocket 使用的 FunASR VAD 与流式 ASR 会话。""" 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: """将浏览器 PCM16 音频转换为 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: """开启一个轮次,并保留短暂的预录音以覆盖起始音素。""" 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]: """将音频送入 VAD,再流式传入当前 ASR 轮次。""" 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: """使用 FunASR 的累计文本创建中间结果事件。""" 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: """刷新 ASR 缓存,然后释放已完成片段的音频。""" 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]: """接收 PCM16 音频并返回 FunASR 中间或最终结果事件。""" 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]: """刷新缓冲音频并清理所有流式缓存。""" 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]