Files
hermes-agent/tools/wake_word_engines.py
T

332 lines
12 KiB
Python

"""Wake-word hotword engines (openWakeWord / sherpa-onnx KWS / Porcupine).
All three run fully on-device. Config, platform probes and sensitivity
accessors live in :mod:`tools.wake_word`; engines read them lazily through
that module so test seams (``patch("tools.wake_word.<name>")``) keep working.
"""
from __future__ import annotations
import logging
import os
from pathlib import Path
from typing import Any, Dict, Optional
logger = logging.getLogger("tools.wake_word")
def _ww():
from tools import wake_word
return wake_word
class _Engine:
"""Minimal hotword-engine contract: feed int16 frames, get a bool."""
frame_length: int = 1280 # 80 ms at 16 kHz
#: (matched phrase, profile name) of the most recent fire. Multi-phrase
#: engines (sherpa) set this for profile routing; single-phrase engines
#: leave it None (callers fall back to configured phrase / active profile).
last_match: Optional[tuple[str, str]] = None
def process(self, frame) -> bool: # frame: 1-D int16 ndarray
raise NotImplementedError
def reset(self) -> None:
"""Clear any internal audio/feature buffer (called on every (re)start)."""
def close(self) -> None:
pass
def _looks_like_path(value: str) -> bool:
return os.sep in value or value.endswith((".onnx", ".tflite", ".ppn")) or os.path.exists(value)
def _sub(cfg: Dict[str, Any], key: str) -> Dict[str, Any]:
sub = cfg.get(key)
return sub if isinstance(sub, dict) else {}
class _OpenWakeWordEngine(_Engine):
"""openWakeWord — free, local ONNX/tflite hotword detection.
Scores one ~80 ms frame at a time; ``sensitivity`` IS the raw 0..1
threshold (higher = stricter). A real utterance holds the score high
across frames while a stray ambient phoneme spikes one, so we require
``confirmation_frames`` consecutive over-threshold frames before firing.
"""
frame_length = 1280 # openWakeWord recommends 80 ms frames.
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.openwakeword", prompt=False)
import openwakeword
from openwakeword.model import Model
ww = _ww()
model_ref = str(_sub(cfg, "openwakeword").get("model") or ww._BUNDLED_MODEL_NAME).strip()
framework = self._usable_framework(ww.resolve_inference_framework(cfg))
self._threshold = ww._sensitivity(cfg)
self._confirm_needed = ww._confirmation_frames(cfg)
self._confirm_streak = 0
# Default (or explicit "hey_hermes") → the bundled model; a built-in
# name or custom path is used as-is.
if model_ref.lower() in ww._BUNDLED_MODEL_ALIASES:
model_ref = ww._bundled_wakeword_path(framework)
# download_models() also fetches the shared feature models (melspectrogram
# + embedding) needed for ANY model, so a custom path must call it too or a
# fresh install crashes on a missing melspectrogram.onnx.
try:
openwakeword.utils.download_models([model_ref])
except Exception as e: # pragma: no cover - network/path dependent
logger.debug("openwakeword model download skipped: %s", e)
self._model = Model(wakeword_models=[model_ref], inference_framework=framework)
self._labels = list(self._model.models.keys())
@staticmethod
def _usable_framework(framework: str) -> str:
"""Refuse openWakeWord's silent tflite→onnx downgrade.
Without a tflite runtime openWakeWord falls back to onnx, which on
macOS ARM64 is the backend whose embedding model never fires — the
listener would arm and stay deaf. Install + bridge the runtime first
(the platform gate lives here because dep specs can't carry PEP 508
markers); on that Mac raise instead of downgrading.
"""
ww = _ww()
if framework != "tflite" or ww.ensure_tflite_runtime():
return framework
try:
from tools import lazy_deps
lazy_deps.ensure("wake.openwakeword.tflite", prompt=False)
except Exception as e:
logger.debug("wake word: tflite runtime install failed: %s", e)
if ww.ensure_tflite_runtime():
return framework
if ww._is_macos_arm64():
raise RuntimeError(
"The wake word needs the tflite backend on this Mac, but its "
"runtime is missing. Install it with: pip install ai-edge-litert"
)
logger.warning("wake word: no tflite runtime available — falling back to onnx")
return "onnx"
def process(self, frame) -> bool:
scores = self._model.predict(frame)
if not any(score >= self._threshold for score in scores.values()):
self._confirm_streak = 0
return False
self._confirm_streak += 1
if self._confirm_streak < self._confirm_needed:
return False
self._confirm_streak = 0
return True
def reset(self) -> None:
# Clears openWakeWord's rolling feature/prediction buffer so stale audio
# captured before a pause can't re-fire the moment we resume.
self._confirm_streak = 0
try:
self._model.reset()
except Exception:
pass
def close(self) -> None:
self.reset()
# sherpa-onnx open-vocabulary KWS model: a small streaming zipformer transducer
# (English, GigaSpeech); one-time download cached under HERMES_HOME. Keywords
# are typed phrases tokenized at RUNTIME — no training step.
_SHERPA_KWS_MODEL_URL = (
"https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/"
"sherpa-onnx-kws-zipformer-gigaspeech-3.3M-2024-01-01.tar.bz2"
)
_SHERPA_KWS_MODEL_DIR = "sherpa-onnx-kws-zipformer-gigaspeech-3.3M-2024-01-01"
def _sherpa_model_root() -> Path:
from hermes_constants import get_hermes_home
return get_hermes_home() / "cache" / "wakewords"
def _ensure_sherpa_model(root: Optional[Path] = None) -> Path:
"""Download + unpack the sherpa KWS model once; return its directory."""
root = root or _sherpa_model_root()
target = root / _SHERPA_KWS_MODEL_DIR
if (target / "tokens.txt").exists():
return target
import tarfile
import urllib.request
root.mkdir(parents=True, exist_ok=True)
archive = root / f"{_SHERPA_KWS_MODEL_DIR}.tar.bz2"
logger.info("wake word: downloading sherpa KWS model (one-time, ~13 MB)")
urllib.request.urlretrieve(_SHERPA_KWS_MODEL_URL, archive) # noqa: S310
with tarfile.open(archive, "r:bz2") as tf:
tf.extractall(root, filter="data")
archive.unlink(missing_ok=True)
if not (target / "tokens.txt").exists():
raise RuntimeError(f"sherpa KWS model unpack failed: {target}")
return target
class _SherpaKwsEngine(_Engine):
"""sherpa-onnx open-vocabulary keyword spotting — any typed phrase, zero training.
``wake_word.phrase`` is BPE-tokenized at runtime against the model's
vocabulary, so here ``phrase`` is DETECTION config, not a cosmetic label.
"""
frame_length = 1280 # streaming zipformer accepts any chunk; match capture path.
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.sherpa", prompt=False)
import sherpa_onnx
from sherpa_onnx import text2token
ww = _ww()
model_dir = str(_sub(cfg, "sherpa").get("model_dir") or "").strip()
d = Path(model_dir) if model_dir else _ensure_sherpa_model()
if not (d / "tokens.txt").exists():
raise RuntimeError(f"sherpa KWS model not found at {d}")
# Phrase set: this profile's own phrase plus — when profile routing is
# on — every other wake-enabled profile's phrase, so ONE listener can
# wake any profile. phrase → profile is kept for routing the match back.
phrase = str(ww._get(cfg, "phrase") or "hey hermes").strip()
phrase_map: Dict[str, str] = {phrase: ww._active_profile_name()}
if bool(cfg.get("profile_routing", True)):
for prof, p in ww.enrolled_profile_phrases().items():
phrase_map.setdefault(p.strip(), prof)
phrases = list(phrase_map)
tokens = text2token(
[p.upper() for p in phrases],
tokens=str(d / "tokens.txt"),
tokens_type="bpe",
bpe_model=str(d / "bpe.model"),
)
import tempfile
# sherpa keyword entries reject spaces in the @display-name; underscore
# them and map display → profile for match routing.
self._display_to_profile: Dict[str, str] = {}
kw = tempfile.NamedTemporaryFile(
mode="w", suffix=".txt", prefix="hermes-kws-", delete=False, encoding="utf-8"
)
for p, toks in zip(phrases, tokens):
display = p.upper().replace(" ", "_")
self._display_to_profile[display] = phrase_map[p]
kw.write(" ".join(toks) + f" @{display}\n")
kw.close()
self._keywords_file = kw.name
self.last_match: Optional[tuple[str, str]] = None
# Shared 0..1 sensitivity → sherpa keywords_threshold. 0.5 lands on
# sherpa's recommended 0.25; a stricter 0.35 missed ~12% of true
# positives in live TTS matrix tests while 0.25 held zero false fires.
threshold = 0.05 + 0.4 * ww._sensitivity(cfg)
def _model_file(pattern: str) -> str:
hits = sorted(d.glob(pattern))
if not hits:
raise RuntimeError(f"sherpa KWS model file missing: {d}/{pattern}")
return str(hits[0])
self._spotter = sherpa_onnx.KeywordSpotter(
tokens=str(d / "tokens.txt"),
encoder=_model_file("encoder-*[!8].onnx"),
decoder=_model_file("decoder-*[!8].onnx"),
joiner=_model_file("joiner-*[!8].onnx"),
keywords_file=self._keywords_file,
keywords_threshold=threshold,
num_threads=1,
)
self._stream = self._spotter.create_stream()
def process(self, frame) -> bool:
import numpy as np
samples = np.asarray(frame, dtype=np.float32) / 32768.0
self._stream.accept_waveform(_ww().SAMPLE_RATE, samples)
fired = False
while self._spotter.is_ready(self._stream):
self._spotter.decode_stream(self._stream)
result = self._spotter.get_result(self._stream)
if result:
fired = True
display = str(result)
self.last_match = (
display.replace("_", " ").lower(),
self._display_to_profile.get(display, ""),
)
# Reset decoder state so one utterance can't fire repeatedly.
self._spotter.reset_stream(self._stream)
return fired
def reset(self) -> None:
# Fresh stream drops buffered audio/decoder state (pause → resume must
# not re-fire on stale audio).
try:
self._stream = self._spotter.create_stream()
except Exception:
pass
def close(self) -> None:
try:
os.unlink(self._keywords_file)
except OSError:
pass
class _PorcupineEngine(_Engine):
"""Picovoice Porcupine — premium, on-device, needs an access key."""
def __init__(self, cfg: Dict[str, Any]):
from tools import lazy_deps
lazy_deps.ensure("wake.porcupine", prompt=False)
import pvporcupine
access_key = (os.getenv("PORCUPINE_ACCESS_KEY") or "").strip()
if not access_key:
raise RuntimeError(
"Porcupine wake word requires PORCUPINE_ACCESS_KEY "
"(get a free key at https://console.picovoice.ai)."
)
keyword = str(_sub(cfg, "porcupine").get("keyword") or "jarvis").strip()
# Porcupine's `sensitivities` runs the OPPOSITE way to our shared knob
# (higher = looser); invert so "higher = stricter" holds for every engine.
kwargs: Dict[str, Any] = {"access_key": access_key, "sensitivities": [1.0 - _ww()._sensitivity(cfg)]}
kwargs["keyword_paths" if _looks_like_path(keyword) else "keywords"] = [keyword]
self._porcupine = pvporcupine.create(**kwargs)
self.frame_length = self._porcupine.frame_length
def process(self, frame) -> bool:
# pvporcupine wants a plain list/sequence of int16 samples.
return self._porcupine.process(frame) >= 0
def close(self) -> None:
try:
self._porcupine.delete()
except Exception:
pass