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()