[Wake/VAD/STT]:完成本地唤醒、人声端点检测与转写入口,包含小杰小杰唤醒、VAD 录音切分、Metadata STT 和 sherpa-onnx 错误边界

This commit is contained in:
mkbk
2026-06-17 18:18:37 +08:00
parent 90ba2bf3b0
commit 4d6232ed29
6 changed files with 389 additions and 5 deletions
+8
View File
@@ -16,6 +16,9 @@ from .models import (
WakeEvent,
)
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
from .wakeword import KeywordWakeWordProvider
from .vad import EnergyVadProvider, VadRecorder
from .stt import MetadataSttProvider, is_valid_transcript_text
__all__ = [
"AppConfig",
@@ -25,6 +28,11 @@ __all__ = [
"ErrorCode",
"FileReplayTransport",
"MemoryAudioTransport",
"KeywordWakeWordProvider",
"EnergyVadProvider",
"VadRecorder",
"MetadataSttProvider",
"is_valid_transcript_text",
"Message",
"PipelineState",
"PlaybackResult",
+93
View File
@@ -0,0 +1,93 @@
from __future__ import annotations
import re
from pathlib import Path
from .models import AudioSegment, ErrorCode, ProviderError, Transcript
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
def is_valid_transcript_text(text: str) -> bool:
return bool(_MEANINGFUL_TEXT.search(text.strip()))
class MetadataSttProvider:
def __init__(self, language: str = "zh") -> None:
self.language = language
self.loaded = False
def load(self) -> None:
self.loaded = True
def transcribe(self, segment: AudioSegment) -> Transcript:
if not self.loaded:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"STT provider is not loaded",
False,
"metadata-stt",
"stt",
)
text = str(segment.metadata.get("transcript", "")).strip()
if not is_valid_transcript_text(text):
raise ProviderError(
ErrorCode.STT_EMPTY_TRANSCRIPT,
"STT produced no meaningful text",
True,
"metadata-stt",
"stt",
)
return Transcript(
text=text,
language=str(segment.metadata.get("language", self.language)),
confidence=float(segment.metadata.get("stt_confidence", 1.0)),
duration_ms=segment.duration_ms,
provider="metadata-stt",
raw_metadata=dict(segment.metadata),
)
class SherpaOnnxSttProvider:
def __init__(self, model_path: str, language: str = "zh") -> None:
self.model_path = Path(model_path)
self.language = language
self.loaded = False
def load(self) -> None:
if not self.model_path.exists():
raise ProviderError(
ErrorCode.STT_MODEL_MISSING,
f"sherpa-onnx STT model path does not exist: {self.model_path}",
False,
"sherpa-onnx-stt",
"stt",
)
try:
import sherpa_onnx # type: ignore[import-not-found] # noqa: F401
except Exception as exc:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
f"sherpa_onnx is not available: {exc}",
False,
"sherpa-onnx-stt",
"stt",
) from exc
self.loaded = True
def transcribe(self, segment: AudioSegment) -> Transcript:
if not self.loaded:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx STT provider is not loaded",
False,
"sherpa-onnx-stt",
"stt",
)
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx runtime transcription adapter requires a concrete model profile",
False,
"sherpa-onnx-stt",
"stt",
)
+123
View File
@@ -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
+53
View File
@@ -0,0 +1,53 @@
from __future__ import annotations
from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent
class KeywordWakeWordProvider:
def __init__(self, keyword: str = "小杰小杰", threshold: float = 0.5) -> None:
self.keyword = keyword
self.threshold = threshold
self.loaded = False
def load(self) -> None:
self.loaded = True
def detect(self, frame: AudioFrame) -> WakeEvent | None:
if not self.loaded:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
"wakeword provider is not loaded",
False,
"keyword-wakeword",
"wakeword",
)
metadata = frame.metadata
confidence = float(metadata.get("wake_confidence", 1.0 if metadata.get("wake") else 0.0))
phrase = str(metadata.get("wake_word", metadata.get("text", "")))
matched = bool(metadata.get("wake")) or phrase.strip() == self.keyword
if matched and confidence >= self.threshold:
return WakeEvent(self.keyword, confidence, frame.timestamp_ms)
return None
def reset(self) -> None:
return None
class MissingWakeWordModelProvider:
def __init__(self, model_path: str) -> None:
self.model_path = model_path
def load(self) -> None:
raise ProviderError(
ErrorCode.WAKE_MODEL_MISSING,
f"wakeword model is missing: {self.model_path}",
False,
"wakeword-model",
"wakeword",
)
def detect(self, frame: AudioFrame) -> WakeEvent | None:
return None
def reset(self) -> None:
return None