Files
Owner/src/owner_voice_pet/wakeword.py
T

179 lines
6.4 KiB
Python

from __future__ import annotations
from pathlib import Path
from typing import Any
from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent
from .speech_models import wake_model_paths
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
class SherpaOnnxKeywordWakeWordProvider:
def __init__(
self,
models_dir: str | Path,
*,
keyword: str = "小杰小杰",
keywords_file: str | Path | None = None,
threshold: float = 0.25,
score: float = 1.0,
sherpa_module: Any | None = None,
) -> None:
self.models_dir = Path(models_dir)
self.keyword = keyword
self.keywords_file = Path(keywords_file) if keywords_file else None
self.threshold = threshold
self.score = score
self.loaded = False
self._sherpa = sherpa_module
self._spotter: Any | None = None
self._stream: Any | None = None
def load(self) -> None:
paths = wake_model_paths(self.models_dir)
if self.keywords_file is not None:
paths["keywords"] = self.keywords_file
missing = [name for name, path in paths.items() if name != "model_dir" and not path.exists()]
if missing:
raise ProviderError(
ErrorCode.WAKE_MODEL_MISSING,
"sherpa-onnx KWS model files are missing: " + ", ".join(missing),
False,
"sherpa-onnx-kws",
"wakeword",
)
if not paths["keywords"].read_text(encoding="utf-8").strip():
raise ProviderError(
ErrorCode.WAKE_MODEL_MISSING,
f"wake keyword file is empty: {paths['keywords']}",
False,
"sherpa-onnx-kws",
"wakeword",
)
sherpa_onnx = self._sherpa
if sherpa_onnx is None:
try:
import sherpa_onnx # type: ignore[import-not-found]
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"sherpa_onnx is not available: {exc}",
False,
"sherpa-onnx-kws",
"wakeword",
) from exc
try:
self._spotter = sherpa_onnx.KeywordSpotter(
tokens=str(paths["tokens"]),
encoder=str(paths["encoder"]),
decoder=str(paths["decoder"]),
joiner=str(paths["joiner"]),
keywords_file=str(paths["keywords"]),
num_threads=2,
sample_rate=16000,
keywords_score=self.score,
keywords_threshold=self.threshold,
provider="cpu",
)
self._stream = self._spotter.create_stream()
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"failed to load sherpa-onnx KWS model: {exc}",
False,
"sherpa-onnx-kws",
"wakeword",
) from exc
self.loaded = True
def detect(self, frame: AudioFrame) -> WakeEvent | None:
if not self.loaded or self._spotter is None or self._stream is None:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
"sherpa-onnx KWS provider is not loaded",
False,
"sherpa-onnx-kws",
"wakeword",
)
try:
import numpy as np
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
if frame.channels > 1 and samples.size:
samples = samples.reshape(-1, frame.channels).mean(axis=1)
self._stream.accept_waveform(frame.sample_rate, samples)
while self._spotter.is_ready(self._stream):
self._spotter.decode_stream(self._stream)
result = str(self._spotter.get_result(self._stream) or "").strip()
if result:
self._spotter.reset_stream(self._stream)
if not self.keyword:
return WakeEvent(result, 1.0, frame.timestamp_ms)
if self.keyword in result or result in self.keyword:
return WakeEvent(self.keyword, 1.0, frame.timestamp_ms)
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"sherpa-onnx KWS detection failed: {exc}",
True,
"sherpa-onnx-kws",
"wakeword",
) from exc
return None
def reset(self) -> None:
if self._spotter is not None and self._stream is not None:
try:
self._spotter.reset_stream(self._stream)
except Exception:
self._stream = self._spotter.create_stream()