392 lines
15 KiB
Python
392 lines
15 KiB
Python
"""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]
|