260 lines
10 KiB
Python
260 lines
10 KiB
Python
"""Wake-word detection via a custom-trained openWakeWord model.
|
|
|
|
Uses `thunderbolt.onnx` — trained specifically for "thunderbolt", the same
|
|
way the main repo's `desk_client/bolt_desk.py` uses `bolt.onnx` for "hey
|
|
bolt". Same runtime (openWakeWord, ONNX inference), same per-frame
|
|
predict()/reset() pattern; the only difference is the model file
|
|
(WAKE_MODEL_FILE) and threshold (WAKE_WORD_THRESHOLD), both configurable via
|
|
.env if a differently-trained model is swapped in later.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections import deque
|
|
from typing import Callable, Optional, Protocol, Union
|
|
|
|
import numpy as np
|
|
|
|
from .. import config
|
|
|
|
|
|
class WakeModel(Protocol):
|
|
def predict(self, frame: np.ndarray) -> dict: ...
|
|
def reset(self) -> None: ...
|
|
|
|
|
|
class NearMissLog:
|
|
"""Rolling record of frames that *almost* fired the wake word.
|
|
|
|
WAKE_WORD_THRESHOLD is otherwise tuned by guessing at a number in .env
|
|
and seeing whether the pet ignores you. Keeping the near misses (scores
|
|
within WAKE_NEAR_MISS_MARGIN below the threshold) turns that into
|
|
evidence: the tray's wake-word tuner shows what your actual "thunderbolt"
|
|
scores, so you can set the threshold just under it.
|
|
|
|
Pure bookkeeping — the caller supplies timestamps, so it's testable.
|
|
"""
|
|
|
|
def __init__(self, limit: int = None, margin: float = None):
|
|
self._entries: deque[tuple[float, float, float]] = deque( # (timestamp, score, threshold)
|
|
maxlen=max(1, config.WAKE_NEAR_MISS_LIMIT if limit is None else limit)
|
|
)
|
|
self._margin = config.WAKE_NEAR_MISS_MARGIN if margin is None else margin
|
|
self._peak = 0.0
|
|
|
|
@property
|
|
def peak(self) -> float:
|
|
"""Highest score seen since the last reset — the "how close did I
|
|
get?" readout while you test the wake phrase."""
|
|
return self._peak
|
|
|
|
def observe(self, score: float, threshold: float, timestamp: float) -> bool:
|
|
"""Record *score*; returns True if it counted as a near miss."""
|
|
self._peak = max(self._peak, score)
|
|
if score >= threshold or score < threshold - self._margin:
|
|
return False
|
|
self._entries.append((timestamp, score, threshold))
|
|
return True
|
|
|
|
def entries(self) -> list[tuple[float, float, float]]:
|
|
return list(self._entries)
|
|
|
|
def clear(self) -> None:
|
|
self._entries.clear()
|
|
self._peak = 0.0
|
|
|
|
|
|
# Where the cached blank state is stashed — on the preprocessor itself
|
|
# rather than in a dict keyed by id(), which CPython reuses after garbage
|
|
# collection and would hand one model another's buffers.
|
|
_BLANK_STATE_ATTR = "_bolt_blank_state"
|
|
|
|
|
|
def _blank_state(preprocessor) -> Optional[tuple]:
|
|
"""(feature_buffer, melspectrogram_buffer) as they are on a freshly
|
|
constructed model — i.e. "having heard nothing". Computing it costs an
|
|
ONNX pass over 10s of silence, so it's cached on the preprocessor and
|
|
copied from thereafter. Returns None if openwakeword's internals don't
|
|
look the way we expect, in which case callers leave the state alone
|
|
rather than corrupting it."""
|
|
cached = getattr(preprocessor, _BLANK_STATE_ATTR, None)
|
|
if cached is not None:
|
|
return cached
|
|
try:
|
|
# Same call the AudioFeatures constructor uses to prime the buffer.
|
|
state = (preprocessor._get_embeddings(np.zeros(160000).astype(np.int16)), np.ones((76, 32)))
|
|
setattr(preprocessor, _BLANK_STATE_ATTR, state)
|
|
except Exception:
|
|
return None
|
|
return state
|
|
|
|
|
|
def hard_reset(model) -> None:
|
|
"""Make the model forget the audio it has already heard — not just its
|
|
predictions.
|
|
|
|
openwakeword's ``Model.reset()`` clears the *prediction* buffer only.
|
|
The rolling audio window the classifier actually scores lives in
|
|
``model.preprocessor`` (raw_data_buffer / melspectrogram_buffer /
|
|
feature_buffer, ~10s of history) and has no reset method of its own. So
|
|
after a detection the wake word is still sitting in that window, and the
|
|
next frame fed to the model re-fires on it — which is exactly what made
|
|
the pet interrupt itself a word into every reply: the "thunderbolt" that
|
|
started the turn was still in the buffer when barge-in resumed feeding it.
|
|
|
|
Never raises. A model whose internals don't match (a fake in tests, a
|
|
future openwakeword release) just gets the plain reset()."""
|
|
try:
|
|
model.reset()
|
|
except Exception:
|
|
pass
|
|
preprocessor = getattr(model, "preprocessor", None)
|
|
if preprocessor is None:
|
|
return
|
|
blank = _blank_state(preprocessor)
|
|
try:
|
|
if blank is not None:
|
|
features, melspectrogram = blank
|
|
preprocessor.feature_buffer = features.copy()
|
|
preprocessor.melspectrogram_buffer = melspectrogram.copy()
|
|
preprocessor.raw_data_buffer.clear()
|
|
preprocessor.accumulated_samples = 0
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def _construct_model(model_cls, model_path: str):
|
|
"""openwakeword's Model() constructor keyword has drifted across
|
|
releases (wakeword_models -> wakeword_model_paths) and some builds
|
|
reject inference_framework entirely — the main repo's ai/wake_word.py
|
|
hit the same drift and works around it the same way: try each known
|
|
calling convention in turn."""
|
|
attempts = [
|
|
lambda: model_cls(wakeword_model_paths=[model_path], inference_framework="onnx"),
|
|
lambda: model_cls(wakeword_model_paths=[model_path]),
|
|
lambda: model_cls(wakeword_models=[model_path], inference_framework="onnx"),
|
|
lambda: model_cls(wakeword_models=[model_path]),
|
|
lambda: model_cls([model_path]),
|
|
]
|
|
last_exc: Optional[TypeError] = None
|
|
for attempt in attempts:
|
|
try:
|
|
return attempt()
|
|
except TypeError as exc:
|
|
last_exc = exc
|
|
raise RuntimeError(
|
|
f"Could not construct openwakeword.Model with any known calling convention "
|
|
f"(last error: {last_exc})"
|
|
)
|
|
|
|
|
|
class _OpenWakeWordModel:
|
|
"""Lazily loads the ONNX model on first use so importing this module
|
|
(and unit-testing listen_for_wake_word with a fake model) never requires
|
|
onnxruntime/openwakeword or the model file to be present."""
|
|
|
|
def __init__(self):
|
|
self._model = None
|
|
|
|
def _ensure_model(self):
|
|
if self._model is None:
|
|
from openwakeword.model import Model
|
|
# from openwakeword.utils import download_models
|
|
|
|
# # The pip package doesn't bundle its melspectrogram/embedding
|
|
# # feature-extraction sub-models — fetch them once on first use
|
|
# # (no-op if already cached in openwakeword's own resources dir).
|
|
# # A non-empty, non-matching model_names list keeps this from
|
|
# # also pulling every official pretrained wakeword model.
|
|
# download_models(model_names=["thunderbolt"])
|
|
|
|
self._model = _construct_model(Model, config.WAKE_MODEL_PATH)
|
|
return self._model
|
|
|
|
@property
|
|
def preprocessor(self):
|
|
"""Proxy the wrapped model's audio-feature buffers.
|
|
|
|
Without this, hard_reset() sees a wrapper with no `preprocessor` and
|
|
silently degrades to openwakeword's shallow reset() — which leaves the
|
|
previous detection sitting in the audio window, i.e. exactly the bug
|
|
hard_reset exists to fix. Returns None before the model is loaded, so
|
|
a reset that happens first is a no-op rather than a load."""
|
|
return getattr(self._model, "preprocessor", None)
|
|
|
|
def predict(self, frame: np.ndarray) -> dict:
|
|
return self._ensure_model().predict(frame)
|
|
|
|
def reset(self) -> None:
|
|
if self._model is not None:
|
|
self._model.reset()
|
|
|
|
|
|
_default_model = _OpenWakeWordModel()
|
|
|
|
|
|
def listen_for_wake_word(
|
|
stream,
|
|
should_continue=lambda: True,
|
|
model: Optional[WakeModel] = None,
|
|
threshold: Union[float, Callable[[], float], None] = None,
|
|
on_tick=None,
|
|
on_score: Optional[Callable[[float, float], None]] = None,
|
|
) -> bool:
|
|
"""Block until the wake word fires (returns True) or *should_continue*
|
|
goes false (returns False).
|
|
|
|
Feeds every frame to *model* (the thunderbolt openWakeWord model by
|
|
default) and treats any class score >= *threshold* as a detection,
|
|
clearing the model's internal state (audio window included, see
|
|
hard_reset) afterward so the next call starts clean — same pattern as
|
|
desk_client/bolt_desk.py's main loop.
|
|
|
|
*threshold* may be a number or a zero-argument callable. The callable
|
|
form exists because this function blocks for minutes at a time: the
|
|
tray's wake-word tuner slider has to be able to change sensitivity
|
|
*during* a listen, not only at the start of the next one.
|
|
|
|
*on_tick*, if given, is called once per ``WAKE_CHECK_INTERVAL_SECONDS``
|
|
(not every frame — prediction is cheap enough to run on every frame, but
|
|
this is the only point control returns to the caller while otherwise
|
|
blocked here for a possibly long time, so it's how a caller drives
|
|
periodic work, e.g. the heartbeat/announcement poll in controller.py,
|
|
during quiet stretches with no wake word).
|
|
|
|
*on_score*, if given, gets ``(best_score, threshold)`` every frame — used
|
|
to log near misses for threshold tuning.
|
|
"""
|
|
model = model or _default_model
|
|
if threshold is None:
|
|
threshold = config.WAKE_WORD_THRESHOLD
|
|
resolve_threshold = threshold if callable(threshold) else (lambda: threshold)
|
|
|
|
frame_len = config.FRAME_LEN
|
|
check_every_frames = max(1, int(config.WAKE_CHECK_INTERVAL_SECONDS * config.SAMPLE_RATE / frame_len))
|
|
frames_since_tick = 0
|
|
|
|
while should_continue():
|
|
chunk, _ = stream.read(frame_len)
|
|
frame = np.asarray(chunk)[:, 0]
|
|
scores = model.predict(frame)
|
|
current_threshold = resolve_threshold()
|
|
best = max(scores.values()) if scores else 0.0
|
|
|
|
frames_since_tick += 1
|
|
if frames_since_tick >= check_every_frames:
|
|
frames_since_tick = 0
|
|
if on_tick is not None:
|
|
on_tick()
|
|
|
|
if on_score is not None:
|
|
on_score(best, current_threshold)
|
|
|
|
if scores and best >= current_threshold:
|
|
# hard_reset, not reset: the phrase has to leave the model's audio
|
|
# window too, or the very next frame we feed it re-fires on the
|
|
# same "thunderbolt" (see hard_reset's docstring).
|
|
hard_reset(model)
|
|
return True
|
|
return False
|