[Wake/VAD/STT]:完成本地唤醒、人声端点检测与转写入口,包含小杰小杰唤醒、VAD 录音切分、Metadata STT 和 sherpa-onnx 错误边界
This commit is contained in:
@@ -0,0 +1,123 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult
|
||||
|
||||
|
||||
class EnergyVadProvider:
|
||||
def __init__(self, threshold: int = 0) -> None:
|
||||
self.threshold = threshold
|
||||
self.loaded = False
|
||||
self._speech_ms = 0
|
||||
self._silence_ms = 0
|
||||
|
||||
def load(self) -> None:
|
||||
self.loaded = True
|
||||
|
||||
def analyze(self, frame: AudioFrame) -> VadResult:
|
||||
if not self.loaded:
|
||||
raise ProviderError(
|
||||
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
||||
"VAD provider is not loaded",
|
||||
False,
|
||||
"energy-vad",
|
||||
"vad",
|
||||
)
|
||||
is_speech = self._is_speech(frame)
|
||||
frame_ms = int(frame.metadata.get("duration_ms", 20))
|
||||
if is_speech:
|
||||
self._speech_ms += frame_ms
|
||||
self._silence_ms = 0
|
||||
else:
|
||||
self._silence_ms += frame_ms
|
||||
return VadResult(
|
||||
is_speech=is_speech,
|
||||
confidence=0.9 if is_speech else 0.1,
|
||||
speech_ms=self._speech_ms,
|
||||
silence_ms=self._silence_ms,
|
||||
)
|
||||
|
||||
def reset(self) -> None:
|
||||
self._speech_ms = 0
|
||||
self._silence_ms = 0
|
||||
|
||||
def _is_speech(self, frame: AudioFrame) -> bool:
|
||||
if "speech" in frame.metadata:
|
||||
return bool(frame.metadata["speech"])
|
||||
if not frame.pcm:
|
||||
return False
|
||||
return any(abs(byte - 128) > self.threshold for byte in frame.pcm)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class VadRecorder:
|
||||
provider: EnergyVadProvider
|
||||
min_duration_ms: int = 300
|
||||
end_silence_ms: int = 200
|
||||
no_speech_timeout_ms: int = 1000
|
||||
max_recording_ms: int = 30000
|
||||
started: bool = field(default=False, init=False)
|
||||
frames: list[AudioFrame] = field(default_factory=list, init=False)
|
||||
first_seen_ms: int | None = field(default=None, init=False)
|
||||
start_time_ms: int | None = field(default=None, init=False)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.reset()
|
||||
|
||||
def reset(self) -> None:
|
||||
self.started = False
|
||||
self.frames: list[AudioFrame] = []
|
||||
self.first_seen_ms: int | None = None
|
||||
self.start_time_ms: int | None = None
|
||||
|
||||
def feed(self, frame: AudioFrame) -> AudioSegment | ProviderError | None:
|
||||
if self.first_seen_ms is None:
|
||||
self.first_seen_ms = frame.timestamp_ms
|
||||
result = self.provider.analyze(frame)
|
||||
if result.is_speech:
|
||||
if not self.started:
|
||||
self.started = True
|
||||
self.start_time_ms = frame.timestamp_ms
|
||||
self.frames.append(frame)
|
||||
elif self.started:
|
||||
self.frames.append(frame)
|
||||
|
||||
if not self.started:
|
||||
elapsed = frame.timestamp_ms - self.first_seen_ms
|
||||
if elapsed >= self.no_speech_timeout_ms:
|
||||
return ProviderError(
|
||||
ErrorCode.VAD_TIMEOUT_NO_SPEECH,
|
||||
"no speech detected after wakeword",
|
||||
True,
|
||||
"energy-vad",
|
||||
"vad",
|
||||
)
|
||||
return None
|
||||
|
||||
start_time = self.start_time_ms if self.start_time_ms is not None else frame.timestamp_ms
|
||||
duration = frame.timestamp_ms - start_time
|
||||
if duration >= self.max_recording_ms:
|
||||
return self._build_segment("max_recording")
|
||||
if result.silence_ms >= self.end_silence_ms and duration >= self.min_duration_ms:
|
||||
return self._build_segment("silence")
|
||||
return None
|
||||
|
||||
def _build_segment(self, end_reason: str) -> AudioSegment:
|
||||
if not self.frames:
|
||||
raise ValueError("cannot build empty segment")
|
||||
metadata: dict[str, object] = {"end_reason": end_reason}
|
||||
for frame in self.frames:
|
||||
metadata.update(dict(frame.metadata))
|
||||
segment = AudioSegment(
|
||||
pcm=b"".join(frame.pcm for frame in self.frames),
|
||||
sample_rate=self.frames[0].sample_rate,
|
||||
channels=self.frames[0].channels,
|
||||
start_time_ms=self.frames[0].timestamp_ms,
|
||||
end_time_ms=self.frames[-1].timestamp_ms
|
||||
+ int(self.frames[-1].metadata.get("duration_ms", 20)),
|
||||
metadata=metadata,
|
||||
)
|
||||
self.reset()
|
||||
self.provider.reset()
|
||||
return segment
|
||||
Reference in New Issue
Block a user