179 lines
6.4 KiB
Python
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()
|