ASR-demo/backend/realtime_websocket/funasr_engine.py

392 lines
15 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""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]