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