Phoneme wake word engine: student+teacher models, INT8 export, C engine, enrollment tooling
Browse files- README.md +221 -0
- act_scales.json +1 -0
- engine_c/pww_decoder.c +169 -0
- engine_c/pww_decoder.h +52 -0
- engine_c/pww_engine.c +254 -0
- engine_c/pww_engine.h +38 -0
- engine_c/pww_frontend.c +127 -0
- engine_c/pww_frontend.h +15 -0
- examples/neptuno_universal.json +69 -0
- frontend_data.h +0 -0
- model_int8.h +0 -0
- phoneme_engine/__init__.py +0 -0
- phoneme_engine/decoder.py +236 -0
- phoneme_engine/enroll.py +254 -0
- phoneme_engine/enroll_universal.py +183 -0
- phoneme_engine/features.py +55 -0
- phoneme_engine/live_demo.py +174 -0
- phoneme_engine/model.py +72 -0
- phoneme_engine/phones.py +83 -0
- phoneme_engine/quantize.py +368 -0
- phoneme_engine/score_phrase.py +340 -0
- phoneme_engine/spot_file.py +76 -0
- phoneme_tcn_student.pt +3 -0
- phoneme_tcn_teacher.pt +3 -0
README.md
ADDED
|
@@ -0,0 +1,221 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
language:
|
| 4 |
+
- en
|
| 5 |
+
tags:
|
| 6 |
+
- keyword-spotting
|
| 7 |
+
- wake-word-detection
|
| 8 |
+
- speech
|
| 9 |
+
- phoneme-recognition
|
| 10 |
+
- ctc
|
| 11 |
+
- esp32
|
| 12 |
+
- tinyml
|
| 13 |
+
- edge-ai
|
| 14 |
+
- on-device
|
| 15 |
+
library_name: pytorch
|
| 16 |
+
pipeline_tag: automatic-speech-recognition
|
| 17 |
+
---
|
| 18 |
+
|
| 19 |
+
# Wake Words Without Training — an open-vocabulary wake word engine for microcontrollers
|
| 20 |
+
|
| 21 |
+
**A wake word here is not a model. It is ~50 bytes of configuration.**
|
| 22 |
+
|
| 23 |
+
One small streaming phoneme recognizer is trained *once* on generic
|
| 24 |
+
speech and never sees a wake word. Any phrase — typed, invented,
|
| 25 |
+
whatever you like — becomes a detector in milliseconds via
|
| 26 |
+
grapheme-to-phoneme lookup, optionally sharpened by five spoken
|
| 27 |
+
examples. The whole runtime fits on an **ESP32-S3** at **340 KB INT8**
|
| 28 |
+
and **14 ms per 40 ms of audio**, leaving wake words changeable
|
| 29 |
+
without retraining, reflashing, or a cloud round-trip.
|
| 30 |
+
|
| 31 |
+
| | Every per-word-trained system | This |
|
| 32 |
+
|---|---|---|
|
| 33 |
+
| New wake word costs | TTS synthesis + GPU training (minutes–hours) | a dictionary lookup (ms) |
|
| 34 |
+
| Artifact per word | a 50–200 KB model | **~50 bytes** (phone ids + threshold) |
|
| 35 |
+
| N simultaneous words | N models | 1 shared model + N tiny decoders |
|
| 36 |
+
| Changing a word on device | reflash | send a few bytes |
|
| 37 |
+
|
| 38 |
+
---
|
| 39 |
+
|
| 40 |
+
## How it works
|
| 41 |
+
|
| 42 |
+
**Runtime lane (wake-word-agnostic, always on):**
|
| 43 |
+
`mic → 40-band log-mel + causal EMA normalization → streaming causal TCN
|
| 44 |
+
→ CTC phoneme posteriors every 20 ms → keyword-filler Viterbi decoder`
|
| 45 |
+
|
| 46 |
+
**Enrollment lane (per word, no training):**
|
| 47 |
+
`text → G2P → phone sequence` and/or `5 spoken examples → phoneme decode
|
| 48 |
+
→ pronunciation variants` → automatic per-variant threshold calibration
|
| 49 |
+
→ **config**, or an explicit *refusal* if the phrase cannot be separated
|
| 50 |
+
from ordinary speech.
|
| 51 |
+
|
| 52 |
+
The decoder is a filler-normalized Viterbi search with four gates, each
|
| 53 |
+
closing a failure mode we hit in practice:
|
| 54 |
+
|
| 55 |
+
1. **Phone-frame normalization** — score is normalized by frames spent in
|
| 56 |
+
*emitting* states only. Normalizing by total duration lets degenerate
|
| 57 |
+
paths idle in blank states (free in silence) and fire rhythmically on
|
| 58 |
+
nothing. Fixing this moved a clean operating point from 65% recall
|
| 59 |
+
@ 12 FA/h to **95% @ 1.25 FA/h**.
|
| 60 |
+
2. **Strong-evidence gate** — ≥50% of phone frames must have their phone
|
| 61 |
+
within 1.5 nats of the frame's best class. Noisy audio yields flat
|
| 62 |
+
posteriors that score fine *on average* while containing nothing;
|
| 63 |
+
this cut false alarms on noise-degraded speech from ~100/h to ~3/h.
|
| 64 |
+
3. **Duration bounds** — 60–200 ms per phone. At a looser cap we observed
|
| 65 |
+
1.5–2.3 s alignments crawling across background television.
|
| 66 |
+
4. **Acoustic energy veto** (on device) — the matched span must exceed the
|
| 67 |
+
rolling noise floor by 5 dB.
|
| 68 |
+
|
| 69 |
+
Detections use only *within-frame logit differences*, so the deployed
|
| 70 |
+
engine never computes a softmax.
|
| 71 |
+
|
| 72 |
+
## Models
|
| 73 |
+
|
| 74 |
+
| File | Params | dev-clean PER | Common Voice dev PER | Use |
|
| 75 |
+
|---|---|---|---|---|
|
| 76 |
+
| `phoneme_tcn_student.pt` | 358,504 | **0.229** | **0.402** | deploy this |
|
| 77 |
+
| `phoneme_tcn_teacher.pt` | 3,337,768 | 0.150 | — | distillation teacher |
|
| 78 |
+
|
| 79 |
+
**Student architecture** — causal TCN: stride-2 stem (k=5), 8
|
| 80 |
+
depthwise-separable causal blocks (k=5, dilations 1,2,4,8 ×2, 192 ch,
|
| 81 |
+
BatchNorm+ReLU, residual), 1×1 head over 40 classes (39 stress-free
|
| 82 |
+
ARPAbet phones + CTC blank). 20 ms output frames, ~2.4 s receptive
|
| 83 |
+
field. Every op is BatchNorm-foldable conv / ReLU / residual — chosen so
|
| 84 |
+
post-training INT8 survives intact (**99.3%** frame-argmax agreement
|
| 85 |
+
with float).
|
| 86 |
+
|
| 87 |
+
**Training** — CTC over phonemized transcripts (CMUdict + neural G2P
|
| 88 |
+
fallback) on **LibriSpeech 960 h + 1.10 M Common Voice 17 English clips**
|
| 89 |
+
(~1,500 h, incl. 83k Indian-accented). On-GPU augmentation: synthetic
|
| 90 |
+
room impulse responses (T60 0.1–0.6 s), additive noise (5–30 dB SNR),
|
| 91 |
+
same-batch babble (10–25 dB), random gain, SpecAugment. The student was
|
| 92 |
+
then distilled from the teacher with **speech-weighted KD** — frames are
|
| 93 |
+
weighted by `1 − p(blank)` because ~75% of CTC frames are blank and
|
| 94 |
+
uniform KD otherwise spends its budget teaching silence (uniform KD
|
| 95 |
+
*degraded* the student; the weighted version improved it).
|
| 96 |
+
|
| 97 |
+
## Benchmarks
|
| 98 |
+
|
| 99 |
+
Ten phrases, Piper LibriTTS-R positives **verified by Whisper** (raw TTS
|
| 100 |
+
is unreliable: "tornado" → *"Pornado"*, "dakota" → *"Decoder"*), against
|
| 101 |
+
1.22 h LibriSpeech dev-clean + 0.82 h Common Voice dev negatives. Noisy
|
| 102 |
+
conditions degrade positives *and* negatives identically so operating
|
| 103 |
+
points stay condition-matched.
|
| 104 |
+
|
| 105 |
+
| Condition | Recall | Notes |
|
| 106 |
+
|---|---|---|
|
| 107 |
+
| Text-only, speaker-independent | median **0.38** @ 0 FA | hardest case: arbitrary phrase, arbitrary speaker, zero examples |
|
| 108 |
+
| Cross-speaker enrollment | repairs dictionary mismatch | tornado 0 → 0.35, "hey jarvis" 0.45 → 0.75 |
|
| 109 |
+
| Personal enrollment (5 examples) | mean **0.554**, median **0.667** @ 3.3 FA/h | synthetic renditions vary more than a self-consistent human |
|
| 110 |
+
| Universal auto-calibration | **83%** of arbitrary voices (`neptuno`) | 36 TTS voices, population-voted variants |
|
| 111 |
+
| On-device, single user | ~**14/15** utterances | ESP32-S3 + INMP441, live session |
|
| 112 |
+
|
| 113 |
+
**Safety property**: calibration *refuses* phrases it cannot separate.
|
| 114 |
+
"norman" (one phone from "normal") was rejected for 5/6 voices and
|
| 115 |
+
"dakota" for 6/6 — independently flagged by the text-only phrase scorer
|
| 116 |
+
(`score_phrase.py`) before any audio existed.
|
| 117 |
+
|
| 118 |
+
## Deployment
|
| 119 |
+
|
| 120 |
+
| Stage | Verification |
|
| 121 |
+
|---|---|
|
| 122 |
+
| BN folding + residual extraction | max err 7.6e-4 vs PyTorch |
|
| 123 |
+
| INT8 simulation vs float | 99.3% frame agreement |
|
| 124 |
+
| **C engine vs INT8 simulation** | **bit-exact (0.0)** |
|
| 125 |
+
| C mel frontend vs training frontend | ≤1e-3 log-mel |
|
| 126 |
+
| ESP32-S3 step time | **14.2 ms** / 40 ms frame (1 core @ 240 MHz) |
|
| 127 |
+
|
| 128 |
+
`engine_c/` is ~300 lines of dependency-free C with one INT8 ring buffer
|
| 129 |
+
per layer, so each frame costs only its own ~350k MACs — no window
|
| 130 |
+
recomputation. On ESP32 the binding constraint was **memory latency, not
|
| 131 |
+
math**: moving weights out of memory-mapped flash (28.6 → 15.9 ms) and
|
| 132 |
+
filling internal SRAM before PSRAM (→ 14.2 ms) mattered far more than
|
| 133 |
+
loop optimization.
|
| 134 |
+
|
| 135 |
+
## Files
|
| 136 |
+
|
| 137 |
+
```
|
| 138 |
+
phoneme_tcn_student.pt deployable model (float, PyTorch state dict)
|
| 139 |
+
phoneme_tcn_teacher.pt distillation teacher
|
| 140 |
+
model_int8.h INT8 weights + layer table (C header, 340 KB)
|
| 141 |
+
act_scales.json activation scales from PTQ calibration
|
| 142 |
+
frontend_data.h mel filterbank + Hann window (exact training values)
|
| 143 |
+
engine_c/ streaming INT8 engine + decoder + mel frontend (C)
|
| 144 |
+
phoneme_engine/ PyTorch model, decoder, enrollment, quantization, scorer
|
| 145 |
+
examples/ a universal wake word config (~50 bytes of JSON)
|
| 146 |
+
```
|
| 147 |
+
|
| 148 |
+
## Usage
|
| 149 |
+
|
| 150 |
+
**Spot a typed phrase (Python):**
|
| 151 |
+
|
| 152 |
+
```python
|
| 153 |
+
import torch
|
| 154 |
+
from phoneme_engine.model import PhonemeTCN
|
| 155 |
+
from phoneme_engine.features import LogMel
|
| 156 |
+
from phoneme_engine.decoder import KeywordSpotter
|
| 157 |
+
|
| 158 |
+
model = PhonemeTCN().eval()
|
| 159 |
+
model.load_state_dict(torch.load("phoneme_tcn_student.pt",
|
| 160 |
+
weights_only=True)["model"])
|
| 161 |
+
frontend = LogMel().eval()
|
| 162 |
+
|
| 163 |
+
spotter = KeywordSpotter("hey orbit") # G2P -> HH EY AO R B AH T
|
| 164 |
+
with torch.no_grad():
|
| 165 |
+
logp = torch.log_softmax(model(frontend(wav)).float(), 2)[0].numpy()
|
| 166 |
+
for frame, score in spotter.run(logp):
|
| 167 |
+
print(f"detected at {frame * 0.02:.2f}s (score {score:.2f})")
|
| 168 |
+
```
|
| 169 |
+
|
| 170 |
+
**Score a candidate phrase before committing to it:**
|
| 171 |
+
|
| 172 |
+
```bash
|
| 173 |
+
python score_phrase.py "vucano"
|
| 174 |
+
# low reliability: /v/ is acoustically weak and often lost on small mics
|
| 175 |
+
```
|
| 176 |
+
|
| 177 |
+
**A wake word config, in full:**
|
| 178 |
+
|
| 179 |
+
```json
|
| 180 |
+
{"phrase": "neptuno",
|
| 181 |
+
"spotters": [{"phones": ["N","EH","P","T","UW","N","OW"],
|
| 182 |
+
"source": "dictionary", "threshold": -2.5}]}
|
| 183 |
+
```
|
| 184 |
+
|
| 185 |
+
**On a microcontroller:** compile `engine_c/` with `model_int8.h` and
|
| 186 |
+
`frontend_data.h`, feed it 40-band mel frames, and pass the logits to
|
| 187 |
+
`pww_spotter_step()`. Wake words are `uint8_t` arrays of phone ids plus a
|
| 188 |
+
float threshold — swap them at runtime.
|
| 189 |
+
|
| 190 |
+
## Limitations
|
| 191 |
+
|
| 192 |
+
- The 358k student is the binding constraint. Text-only speaker-independent
|
| 193 |
+
recall is modest; 40% PER on real-world speech leaves little margin for
|
| 194 |
+
accented, distant, or noisy input. Personal enrollment recovers much of it.
|
| 195 |
+
- Benchmarks are TTS-based (Whisper-verified, but synthetic). Human
|
| 196 |
+
multi-speaker evaluation is future work.
|
| 197 |
+
- **Onset phonetics dominate word quality.** Nasals/stops/sibilants
|
| 198 |
+
(N, M, K, T, S) work well; weak fricatives (V, F, TH, H) are often lost
|
| 199 |
+
on small MEMS mics. The included scorer predicts this from text.
|
| 200 |
+
- False-alarm rates are measured against *continuous speech* — a worst case
|
| 201 |
+
versus mostly-quiet rooms.
|
| 202 |
+
- Single-microphone. Far-field performance is a hardware question
|
| 203 |
+
(beamforming arrays), not a decoder one.
|
| 204 |
+
- English only, though nothing in the architecture is language-bound beyond
|
| 205 |
+
the phone inventory and G2P.
|
| 206 |
+
|
| 207 |
+
## License
|
| 208 |
+
|
| 209 |
+
Apache-2.0. Trained on LibriSpeech (CC BY 4.0) and Common Voice 17 (CC0).
|
| 210 |
+
|
| 211 |
+
## Citation
|
| 212 |
+
|
| 213 |
+
```bibtex
|
| 214 |
+
@misc{wakewordswithouttraining2026,
|
| 215 |
+
title = {Wake Words Without Training: Open-Vocabulary Wake Word
|
| 216 |
+
Creation from Text and a Few Examples},
|
| 217 |
+
author = {IOTEverythin},
|
| 218 |
+
year = {2026},
|
| 219 |
+
url = {https://huggingface.co/IOTEverythin/phoneme-wake-word}
|
| 220 |
+
}
|
| 221 |
+
```
|
act_scales.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
[0.11520339935783326, 0.023926198496594837, 0.12479096896653497, 0.040990088034201386, 0.1265573037491855, 0.07117628648404385, 0.12578543868460848, 0.07480658649596755, 0.12323264005650608, 0.09615841090999173, 0.13713947542410823, 0.1672581903290684, 0.11631624975349672, 0.17186096197009412, 0.09403160310875937, 0.16450792622578664, 0.0586112993594116, 0.16962655659424575, 0.30934620578167427]
|
engine_c/pww_decoder.c
ADDED
|
@@ -0,0 +1,169 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// C port of KeywordSpotter (see phoneme_engine/decoder.py for the
|
| 2 |
+
// reference semantics and design rationale).
|
| 3 |
+
#include "pww_decoder.h"
|
| 4 |
+
|
| 5 |
+
#include <string.h>
|
| 6 |
+
|
| 7 |
+
#include "model_int8.h"
|
| 8 |
+
|
| 9 |
+
#define NEG_INF (-1e30f)
|
| 10 |
+
// 200 ms/phone cap: looser caps let garbage alignments crawl across
|
| 11 |
+
// continuous background speech (see decoder.py)
|
| 12 |
+
#define MAX_FRAMES_PER_PHONE 10
|
| 13 |
+
|
| 14 |
+
// Confusable sets, mirroring decoder.py CONFUSABLE. Ids are 1-based
|
| 15 |
+
// (0 = CTC blank), matching PWW_PHONES order in model_int8.h:
|
| 16 |
+
// AA=1 AE=2 AH=3 AO=4 AW=5 AY=6 B=7 CH=8 D=9 DH=10 EH=11 ER=12 EY=13
|
| 17 |
+
// F=14 G=15 HH=16 IH=17 IY=18 JH=19 K=20 L=21 M=22 N=23 NG=24 OW=25
|
| 18 |
+
// OY=26 P=27 R=28 S=29 SH=30 T=31 TH=32 UH=33 UW=34 V=35 W=36 Y=37
|
| 19 |
+
// Z=38 ZH=39
|
| 20 |
+
static int confusable(uint8_t id, uint8_t out[PWW_DEC_MAX_ALLOWED]) {
|
| 21 |
+
switch (id) {
|
| 22 |
+
case 3: out[0] = 3; out[1] = 17; out[2] = 12; return 3; // AH
|
| 23 |
+
case 17: out[0] = 17; out[1] = 3; return 2; // IH
|
| 24 |
+
case 12: out[0] = 12; out[1] = 3; return 2; // ER
|
| 25 |
+
case 4: out[0] = 4; out[1] = 1; return 2; // AO
|
| 26 |
+
case 2: out[0] = 2; out[1] = 1; return 2; // AE
|
| 27 |
+
case 33: out[0] = 33; out[1] = 34; return 2; // UH
|
| 28 |
+
case 20: out[0] = 20; out[1] = 15; return 2; // K
|
| 29 |
+
case 15: out[0] = 15; out[1] = 20; return 2; // G
|
| 30 |
+
case 31: out[0] = 31; out[1] = 9; return 2; // T
|
| 31 |
+
case 9: out[0] = 9; out[1] = 31; return 2; // D
|
| 32 |
+
default: out[0] = id; return 1;
|
| 33 |
+
}
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
int pww_spotter_init(pww_spotter_t *sp, const uint8_t *phone_ids,
|
| 37 |
+
int n_phones, float threshold) {
|
| 38 |
+
if (n_phones < 2 || n_phones > PWW_DEC_MAX_PHONES) return -1;
|
| 39 |
+
memset(sp, 0, sizeof(*sp));
|
| 40 |
+
sp->n_phones = n_phones;
|
| 41 |
+
sp->threshold = threshold;
|
| 42 |
+
sp->strong_margin = 1.5f;
|
| 43 |
+
sp->strong_ratio = 0.5f;
|
| 44 |
+
sp->refractory = 50;
|
| 45 |
+
|
| 46 |
+
int n = 0;
|
| 47 |
+
// leading blank
|
| 48 |
+
sp->labels[n] = 0;
|
| 49 |
+
sp->preds[n][0] = 0; sp->preds[n][1] = -1;
|
| 50 |
+
sp->entry[n] = 1;
|
| 51 |
+
n++;
|
| 52 |
+
for (int i = 0; i < n_phones; i++) {
|
| 53 |
+
uint8_t pid = phone_ids[i];
|
| 54 |
+
int a = n;
|
| 55 |
+
if (i == 0) {
|
| 56 |
+
sp->labels[n] = pid;
|
| 57 |
+
sp->preds[n][0] = 0; sp->preds[n][1] = -1;
|
| 58 |
+
sp->entry[n] = 1;
|
| 59 |
+
n++;
|
| 60 |
+
} else {
|
| 61 |
+
sp->labels[n] = pid;
|
| 62 |
+
sp->preds[n][0] = (int8_t)(a - 1);
|
| 63 |
+
// skip the blank between different phones
|
| 64 |
+
sp->preds[n][1] = (sp->labels[a - 2] != pid)
|
| 65 |
+
? (int8_t)(a - 2) : -1;
|
| 66 |
+
sp->entry[n] = 0;
|
| 67 |
+
n++;
|
| 68 |
+
}
|
| 69 |
+
sp->labels[n] = pid; // B state, only from A
|
| 70 |
+
sp->preds[n][0] = (int8_t)a; sp->preds[n][1] = -1;
|
| 71 |
+
sp->entry[n] = 0;
|
| 72 |
+
n++;
|
| 73 |
+
sp->labels[n] = 0; // blank after phone
|
| 74 |
+
sp->preds[n][0] = (int8_t)(a + 1); sp->preds[n][1] = -1;
|
| 75 |
+
sp->entry[n] = 0;
|
| 76 |
+
n++;
|
| 77 |
+
}
|
| 78 |
+
sp->n_states = n;
|
| 79 |
+
for (int s = 0; s < n; s++)
|
| 80 |
+
sp->n_allowed[s] = (sp->labels[s] == 0)
|
| 81 |
+
? (sp->allowed[s][0] = 0, 1)
|
| 82 |
+
: (uint8_t)confusable(sp->labels[s], sp->allowed[s]);
|
| 83 |
+
sp->min_dur = 3 * n_phones;
|
| 84 |
+
sp->max_dur = MAX_FRAMES_PER_PHONE * n_phones;
|
| 85 |
+
pww_spotter_reset(sp);
|
| 86 |
+
return 0;
|
| 87 |
+
}
|
| 88 |
+
|
| 89 |
+
void pww_spotter_reset(pww_spotter_t *sp) {
|
| 90 |
+
for (int s = 0; s < sp->n_states; s++) {
|
| 91 |
+
sp->rel[s] = NEG_INF;
|
| 92 |
+
sp->start[s] = 0;
|
| 93 |
+
sp->pframes[s] = 0;
|
| 94 |
+
sp->strong[s] = 0;
|
| 95 |
+
}
|
| 96 |
+
sp->cooldown = 0;
|
| 97 |
+
sp->t = 0;
|
| 98 |
+
sp->best_seen = NEG_INF;
|
| 99 |
+
}
|
| 100 |
+
|
| 101 |
+
int pww_spotter_step(pww_spotter_t *sp, const float *logits,
|
| 102 |
+
float *score_out, int *dur_out) {
|
| 103 |
+
float filler = logits[0];
|
| 104 |
+
for (int c = 1; c < PWW_NUM_CLASSES; c++)
|
| 105 |
+
if (logits[c] > filler) filler = logits[c];
|
| 106 |
+
|
| 107 |
+
float rel[PWW_DEC_MAX_STATES];
|
| 108 |
+
int32_t start[PWW_DEC_MAX_STATES], pf[PWW_DEC_MAX_STATES],
|
| 109 |
+
strong[PWW_DEC_MAX_STATES];
|
| 110 |
+
|
| 111 |
+
for (int s = 0; s < sp->n_states; s++) {
|
| 112 |
+
float best = sp->rel[s];
|
| 113 |
+
int32_t b_start = sp->start[s], b_pf = sp->pframes[s],
|
| 114 |
+
b_sf = sp->strong[s];
|
| 115 |
+
for (int q = 0; q < 2; q++) {
|
| 116 |
+
int8_t p = sp->preds[s][q];
|
| 117 |
+
if (p >= 0 && sp->rel[p] > best) {
|
| 118 |
+
best = sp->rel[p];
|
| 119 |
+
b_start = sp->start[p];
|
| 120 |
+
b_pf = sp->pframes[p];
|
| 121 |
+
b_sf = sp->strong[p];
|
| 122 |
+
}
|
| 123 |
+
}
|
| 124 |
+
if (sp->entry[s] && 0.0f >= best) {
|
| 125 |
+
best = 0.0f; b_start = sp->t; b_pf = 0; b_sf = 0;
|
| 126 |
+
}
|
| 127 |
+
float emit = logits[sp->allowed[s][0]];
|
| 128 |
+
for (int a = 1; a < sp->n_allowed[s]; a++) {
|
| 129 |
+
float v = logits[sp->allowed[s][a]];
|
| 130 |
+
if (v > emit) emit = v;
|
| 131 |
+
}
|
| 132 |
+
int is_phone = sp->labels[s] != 0;
|
| 133 |
+
rel[s] = best + emit - filler;
|
| 134 |
+
start[s] = b_start;
|
| 135 |
+
pf[s] = b_pf + (is_phone ? 1 : 0);
|
| 136 |
+
strong[s] = b_sf +
|
| 137 |
+
((is_phone && emit >= filler - sp->strong_margin) ? 1 : 0);
|
| 138 |
+
}
|
| 139 |
+
memcpy(sp->rel, rel, sizeof(float) * sp->n_states);
|
| 140 |
+
memcpy(sp->start, start, sizeof(int32_t) * sp->n_states);
|
| 141 |
+
memcpy(sp->pframes, pf, sizeof(int32_t) * sp->n_states);
|
| 142 |
+
memcpy(sp->strong, strong, sizeof(int32_t) * sp->n_states);
|
| 143 |
+
sp->t++;
|
| 144 |
+
|
| 145 |
+
float norm_best = NEG_INF;
|
| 146 |
+
int norm_dur = 0, have = 0;
|
| 147 |
+
for (int f = 0; f < 2; f++) {
|
| 148 |
+
int s = sp->n_states - 1 - f;
|
| 149 |
+
int32_t dur = sp->t - sp->start[s];
|
| 150 |
+
int32_t p = sp->pframes[s];
|
| 151 |
+
if (dur < sp->min_dur || dur > sp->max_dur) continue;
|
| 152 |
+
if (p < 2 * sp->n_phones) continue;
|
| 153 |
+
if ((float)sp->strong[s] < sp->strong_ratio * (float)p) continue;
|
| 154 |
+
float norm = sp->rel[s] / (float)p;
|
| 155 |
+
if (!have || norm > norm_best) {
|
| 156 |
+
norm_best = norm; norm_dur = (int)dur; have = 1;
|
| 157 |
+
}
|
| 158 |
+
}
|
| 159 |
+
if (have && norm_best > sp->best_seen) sp->best_seen = norm_best;
|
| 160 |
+
if (sp->cooldown > 0) { sp->cooldown--; return 0; }
|
| 161 |
+
if (have && norm_best > sp->threshold) {
|
| 162 |
+
sp->cooldown = sp->refractory;
|
| 163 |
+
for (int s = 0; s < sp->n_states; s++) sp->rel[s] = NEG_INF;
|
| 164 |
+
*score_out = norm_best;
|
| 165 |
+
*dur_out = norm_dur;
|
| 166 |
+
return 1;
|
| 167 |
+
}
|
| 168 |
+
return 0;
|
| 169 |
+
}
|
engine_c/pww_decoder.h
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Keyword-filler Viterbi spotter over streaming logits - C port of
|
| 2 |
+
// phoneme_engine/decoder.py (KeywordSpotter). Same states, gates and
|
| 3 |
+
// scoring; consumes raw logits (only within-frame differences matter).
|
| 4 |
+
#pragma once
|
| 5 |
+
#include <stdint.h>
|
| 6 |
+
|
| 7 |
+
#ifdef __cplusplus
|
| 8 |
+
extern "C" {
|
| 9 |
+
#endif
|
| 10 |
+
|
| 11 |
+
#define PWW_DEC_MAX_PHONES 24
|
| 12 |
+
#define PWW_DEC_MAX_STATES (3 * PWW_DEC_MAX_PHONES + 1)
|
| 13 |
+
#define PWW_DEC_MAX_ALLOWED 3
|
| 14 |
+
|
| 15 |
+
typedef struct {
|
| 16 |
+
int n_phones;
|
| 17 |
+
int n_states;
|
| 18 |
+
uint8_t labels[PWW_DEC_MAX_STATES]; // class id per state
|
| 19 |
+
uint8_t n_allowed[PWW_DEC_MAX_STATES];
|
| 20 |
+
uint8_t allowed[PWW_DEC_MAX_STATES][PWW_DEC_MAX_ALLOWED];
|
| 21 |
+
int8_t preds[PWW_DEC_MAX_STATES][2]; // predecessor ids, -1 unused
|
| 22 |
+
uint8_t entry[PWW_DEC_MAX_STATES];
|
| 23 |
+
int min_dur, max_dur;
|
| 24 |
+
float threshold;
|
| 25 |
+
float strong_margin; // 1.5
|
| 26 |
+
float strong_ratio; // 0.5
|
| 27 |
+
int refractory; // frames
|
| 28 |
+
|
| 29 |
+
// runtime state
|
| 30 |
+
float rel[PWW_DEC_MAX_STATES];
|
| 31 |
+
int32_t start[PWW_DEC_MAX_STATES];
|
| 32 |
+
int32_t pframes[PWW_DEC_MAX_STATES];
|
| 33 |
+
int32_t strong[PWW_DEC_MAX_STATES];
|
| 34 |
+
int cooldown;
|
| 35 |
+
int32_t t;
|
| 36 |
+
float best_seen; // best gate-passing score since last clear (diag)
|
| 37 |
+
} pww_spotter_t;
|
| 38 |
+
|
| 39 |
+
// phone_ids: sequence of class ids (1..39, from PWW_PHONES order +1).
|
| 40 |
+
// Returns 0 on success, -1 if too long/short.
|
| 41 |
+
int pww_spotter_init(pww_spotter_t *sp, const uint8_t *phone_ids,
|
| 42 |
+
int n_phones, float threshold);
|
| 43 |
+
void pww_spotter_reset(pww_spotter_t *sp);
|
| 44 |
+
|
| 45 |
+
// One frame of logits (PWW_NUM_CLASSES floats). Returns 1 and fills
|
| 46 |
+
// score/dur when the keyword fires, else 0.
|
| 47 |
+
int pww_spotter_step(pww_spotter_t *sp, const float *logits,
|
| 48 |
+
float *score_out, int *dur_out);
|
| 49 |
+
|
| 50 |
+
#ifdef __cplusplus
|
| 51 |
+
}
|
| 52 |
+
#endif
|
engine_c/pww_engine.c
ADDED
|
@@ -0,0 +1,254 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Streaming int8 phoneme engine. See pww_engine.h for the contract.
|
| 2 |
+
#include "pww_engine.h"
|
| 3 |
+
|
| 4 |
+
#include <math.h>
|
| 5 |
+
#include <stdlib.h>
|
| 6 |
+
#include <string.h>
|
| 7 |
+
|
| 8 |
+
#include "model_int8.h"
|
| 9 |
+
|
| 10 |
+
// Debug: when >= 0, pww_engine_step returns layer N's snapped int8
|
| 11 |
+
// output column instead of the final logits.
|
| 12 |
+
int pww_dump_layer = -1;
|
| 13 |
+
|
| 14 |
+
// Ring buffer per layer holding the int8 input history a causal conv
|
| 15 |
+
// needs: (k-1)*dilation + 1 columns of in_c values.
|
| 16 |
+
typedef struct {
|
| 17 |
+
int8_t *buf; // hist * in_c, column-major by time step
|
| 18 |
+
int hist; // number of columns
|
| 19 |
+
int pos; // next write slot
|
| 20 |
+
int primed; // columns written so far (zeros before that)
|
| 21 |
+
} ring_t;
|
| 22 |
+
|
| 23 |
+
struct pww_engine {
|
| 24 |
+
ring_t rings[PWW_NUM_LAYERS];
|
| 25 |
+
// weights copied out of memory-mapped flash at create() time: flash
|
| 26 |
+
// cache misses were 20x slower than the arithmetic
|
| 27 |
+
pww_layer_t layers[PWW_NUM_LAYERS];
|
| 28 |
+
// block input snapshot (int8 col + its scale) for residual adds
|
| 29 |
+
int8_t block_in[512];
|
| 30 |
+
float block_in_scale;
|
| 31 |
+
float scratch_f[512];
|
| 32 |
+
int8_t col_a[512], col_b[512];
|
| 33 |
+
};
|
| 34 |
+
|
| 35 |
+
// Weight allocation. On ESP32 the engine is memory-latency bound, not
|
| 36 |
+
// MAC bound: internal SRAM is several times faster to walk than octal
|
| 37 |
+
// PSRAM, so fill the internal budget first and spill the rest to PSRAM.
|
| 38 |
+
#ifdef ESP_PLATFORM
|
| 39 |
+
#include "esp_heap_caps.h"
|
| 40 |
+
#ifndef PWW_INTERNAL_WEIGHT_BUDGET
|
| 41 |
+
#define PWW_INTERNAL_WEIGHT_BUDGET (160 * 1024)
|
| 42 |
+
#endif
|
| 43 |
+
static size_t s_internal_used = 0;
|
| 44 |
+
static void *weights_alloc(size_t n) {
|
| 45 |
+
if (s_internal_used + n <= PWW_INTERNAL_WEIGHT_BUDGET) {
|
| 46 |
+
void *p = heap_caps_malloc(n, MALLOC_CAP_INTERNAL | MALLOC_CAP_8BIT);
|
| 47 |
+
if (p) { s_internal_used += n; return p; }
|
| 48 |
+
}
|
| 49 |
+
void *p = heap_caps_malloc(n, MALLOC_CAP_SPIRAM);
|
| 50 |
+
if (!p) p = malloc(n);
|
| 51 |
+
return p;
|
| 52 |
+
}
|
| 53 |
+
size_t pww_internal_weight_bytes(void) { return s_internal_used; }
|
| 54 |
+
// engine state (ring buffers, scratch) is touched every frame - never
|
| 55 |
+
// let it land in PSRAM
|
| 56 |
+
static void *state_alloc(size_t n) {
|
| 57 |
+
void *p = heap_caps_calloc(1, n, MALLOC_CAP_INTERNAL | MALLOC_CAP_8BIT);
|
| 58 |
+
return p ? p : calloc(1, n);
|
| 59 |
+
}
|
| 60 |
+
#else
|
| 61 |
+
static void *weights_alloc(size_t n) { return malloc(n); }
|
| 62 |
+
static void *state_alloc(size_t n) { return calloc(1, n); }
|
| 63 |
+
#endif
|
| 64 |
+
|
| 65 |
+
static int8_t quant_clamp(float v, float inv_scale) {
|
| 66 |
+
float q = roundf(v * inv_scale);
|
| 67 |
+
if (q > 127.f) q = 127.f;
|
| 68 |
+
if (q < -127.f) q = -127.f;
|
| 69 |
+
return (int8_t)q;
|
| 70 |
+
}
|
| 71 |
+
|
| 72 |
+
pww_engine_t *pww_engine_create(void) {
|
| 73 |
+
pww_engine_t *e = (pww_engine_t *)state_alloc(sizeof(pww_engine_t));
|
| 74 |
+
if (!e) return NULL;
|
| 75 |
+
for (int li = 0; li < PWW_NUM_LAYERS; li++) {
|
| 76 |
+
e->layers[li] = PWW_LAYERS[li];
|
| 77 |
+
pww_layer_t *M = &e->layers[li];
|
| 78 |
+
size_t wn = (size_t)(M->is_dw ? M->out_c : M->in_c * M->out_c)
|
| 79 |
+
* M->k;
|
| 80 |
+
int8_t *wcopy = (int8_t *)weights_alloc(wn);
|
| 81 |
+
float *scopy = (float *)malloc(sizeof(float) * M->out_c);
|
| 82 |
+
float *bcopy = (float *)malloc(sizeof(float) * M->out_c);
|
| 83 |
+
if (!wcopy || !scopy || !bcopy) { pww_engine_destroy(e); return NULL; }
|
| 84 |
+
memcpy(wcopy, M->w, wn);
|
| 85 |
+
memcpy(scopy, M->s, sizeof(float) * M->out_c);
|
| 86 |
+
memcpy(bcopy, M->b, sizeof(float) * M->out_c);
|
| 87 |
+
M->w = wcopy; M->s = scopy; M->b = bcopy;
|
| 88 |
+
const pww_layer_t *L = &PWW_LAYERS[li];
|
| 89 |
+
int hist = (L->k - 1) * L->dilation + 1;
|
| 90 |
+
// stride-2 stem consumes 2 input columns per step
|
| 91 |
+
if (L->stride == 2) hist += 1;
|
| 92 |
+
ring_t *r = &e->rings[li];
|
| 93 |
+
r->hist = hist;
|
| 94 |
+
r->buf = (int8_t *)state_alloc((size_t)hist * L->in_c);
|
| 95 |
+
if (!r->buf) { pww_engine_destroy(e); return NULL; }
|
| 96 |
+
}
|
| 97 |
+
pww_engine_reset(e);
|
| 98 |
+
return e;
|
| 99 |
+
}
|
| 100 |
+
|
| 101 |
+
void pww_engine_destroy(pww_engine_t *e) {
|
| 102 |
+
if (!e) return;
|
| 103 |
+
for (int li = 0; li < PWW_NUM_LAYERS; li++) {
|
| 104 |
+
free(e->rings[li].buf);
|
| 105 |
+
if (e->layers[li].w && e->layers[li].w != PWW_LAYERS[li].w) {
|
| 106 |
+
free((void *)e->layers[li].w);
|
| 107 |
+
free((void *)e->layers[li].s);
|
| 108 |
+
free((void *)e->layers[li].b);
|
| 109 |
+
}
|
| 110 |
+
}
|
| 111 |
+
free(e);
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
void pww_engine_reset(pww_engine_t *e) {
|
| 115 |
+
for (int li = 0; li < PWW_NUM_LAYERS; li++) {
|
| 116 |
+
ring_t *r = &e->rings[li];
|
| 117 |
+
memset(r->buf, 0, (size_t)r->hist * PWW_LAYERS[li].in_c);
|
| 118 |
+
r->pos = 0;
|
| 119 |
+
r->primed = 0;
|
| 120 |
+
}
|
| 121 |
+
}
|
| 122 |
+
|
| 123 |
+
static void ring_push(ring_t *r, const int8_t *col, int in_c) {
|
| 124 |
+
memcpy(r->buf + (size_t)r->pos * in_c, col, (size_t)in_c);
|
| 125 |
+
r->pos = (r->pos + 1) % r->hist;
|
| 126 |
+
if (r->primed < r->hist) r->primed++;
|
| 127 |
+
}
|
| 128 |
+
|
| 129 |
+
// column at "delay" steps in the past (0 = newest)
|
| 130 |
+
static const int8_t *ring_at(const ring_t *r, int delay, int in_c) {
|
| 131 |
+
int idx = r->pos - 1 - delay;
|
| 132 |
+
while (idx < 0) idx += r->hist;
|
| 133 |
+
return r->buf + (size_t)idx * in_c;
|
| 134 |
+
}
|
| 135 |
+
|
| 136 |
+
// Runs one layer on the newest ring content, writing float pre-snap
|
| 137 |
+
// output to out_f (out_c values). For stride-2 layers the newest TWO
|
| 138 |
+
// columns have been pushed before calling.
|
| 139 |
+
static void layer_forward(const pww_layer_t *L, const ring_t *r,
|
| 140 |
+
float *out_f) {
|
| 141 |
+
int k = L->k, d = L->dilation, in_c = L->in_c, out_c = L->out_c;
|
| 142 |
+
// stride-2 layers get two pushes per step but output frame t only
|
| 143 |
+
// consumes up to input column 2t: the newest tap sits one column back
|
| 144 |
+
int off = (L->stride == 2) ? 1 : 0;
|
| 145 |
+
if (L->is_dw) {
|
| 146 |
+
for (int c = 0; c < out_c; c++) out_f[c] = 0.f;
|
| 147 |
+
for (int i = 0; i < k; i++) {
|
| 148 |
+
// tap i is the newest at i == k-1
|
| 149 |
+
const int8_t *col = ring_at(r, (k - 1 - i) * d + off, in_c);
|
| 150 |
+
const int8_t *w = L->w + i; // w layout: (ch, 1, k)
|
| 151 |
+
for (int c = 0; c < out_c; c++)
|
| 152 |
+
out_f[c] += (float)((int32_t)col[c] * (int32_t)w[c * k]);
|
| 153 |
+
}
|
| 154 |
+
for (int c = 0; c < out_c; c++)
|
| 155 |
+
out_f[c] = out_f[c] * L->s[c] + L->b[c];
|
| 156 |
+
} else if (k == 1) {
|
| 157 |
+
// 1x1 conv = matrix-vector; 85% of all MACs live here.
|
| 158 |
+
// Contiguous rows, 4-way unroll, single ring lookup.
|
| 159 |
+
const int8_t *restrict col = ring_at(r, off, in_c);
|
| 160 |
+
for (int c = 0; c < out_c; c++) {
|
| 161 |
+
const int8_t *restrict w = L->w + (size_t)c * in_c;
|
| 162 |
+
int32_t a0 = 0, a1 = 0, a2 = 0, a3 = 0;
|
| 163 |
+
int j = 0;
|
| 164 |
+
for (; j + 4 <= in_c; j += 4) {
|
| 165 |
+
a0 += (int32_t)col[j] * (int32_t)w[j];
|
| 166 |
+
a1 += (int32_t)col[j + 1] * (int32_t)w[j + 1];
|
| 167 |
+
a2 += (int32_t)col[j + 2] * (int32_t)w[j + 2];
|
| 168 |
+
a3 += (int32_t)col[j + 3] * (int32_t)w[j + 3];
|
| 169 |
+
}
|
| 170 |
+
int32_t acc = a0 + a1 + a2 + a3;
|
| 171 |
+
for (; j < in_c; j++)
|
| 172 |
+
acc += (int32_t)col[j] * (int32_t)w[j];
|
| 173 |
+
out_f[c] = (float)acc * L->s[c] + L->b[c];
|
| 174 |
+
}
|
| 175 |
+
} else {
|
| 176 |
+
// general conv: hoist the k column pointers out of the c loop
|
| 177 |
+
const int8_t *cols[8];
|
| 178 |
+
for (int i = 0; i < k; i++)
|
| 179 |
+
cols[i] = ring_at(r, (k - 1 - i) * d + off, in_c);
|
| 180 |
+
for (int c = 0; c < out_c; c++) {
|
| 181 |
+
const int8_t *restrict w = L->w + (size_t)c * in_c * k;
|
| 182 |
+
int32_t acc = 0;
|
| 183 |
+
for (int i = 0; i < k; i++) {
|
| 184 |
+
const int8_t *restrict col = cols[i];
|
| 185 |
+
const int8_t *restrict wk = w + i;
|
| 186 |
+
for (int j = 0; j < in_c; j++)
|
| 187 |
+
acc += (int32_t)col[j] * (int32_t)wk[(size_t)j * k];
|
| 188 |
+
}
|
| 189 |
+
out_f[c] = (float)acc * L->s[c] + L->b[c];
|
| 190 |
+
}
|
| 191 |
+
}
|
| 192 |
+
}
|
| 193 |
+
|
| 194 |
+
int pww_engine_step(pww_engine_t *e, const float *mel0, const float *mel1,
|
| 195 |
+
float *logits_out) {
|
| 196 |
+
// quantize the two input mel columns to the input scale
|
| 197 |
+
float inv_in = 1.0f / PWW_INPUT_SCALE;
|
| 198 |
+
for (int j = 0; j < PWW_MELS; j++)
|
| 199 |
+
e->col_a[j] = quant_clamp(mel0[j], inv_in);
|
| 200 |
+
ring_push(&e->rings[0], e->col_a, PWW_MELS);
|
| 201 |
+
for (int j = 0; j < PWW_MELS; j++)
|
| 202 |
+
e->col_a[j] = quant_clamp(mel1[j], inv_in);
|
| 203 |
+
ring_push(&e->rings[0], e->col_a, PWW_MELS);
|
| 204 |
+
|
| 205 |
+
int8_t *cur = e->col_a; // int8 column flowing between layers
|
| 206 |
+
float cur_scale = PWW_INPUT_SCALE;
|
| 207 |
+
(void)cur_scale;
|
| 208 |
+
|
| 209 |
+
for (int li = 0; li < PWW_NUM_LAYERS; li++) {
|
| 210 |
+
const pww_layer_t *L = &e->layers[li];
|
| 211 |
+
ring_t *r = &e->rings[li];
|
| 212 |
+
if (li > 0) ring_push(r, cur, L->in_c);
|
| 213 |
+
|
| 214 |
+
if (L->is_dw) {
|
| 215 |
+
// save the block input column + scale for the residual 2
|
| 216 |
+
// layers later (dw -> pw(residual))
|
| 217 |
+
memcpy(e->block_in, ring_at(r, 0, L->in_c), (size_t)L->in_c);
|
| 218 |
+
e->block_in_scale =
|
| 219 |
+
(li == 0) ? PWW_INPUT_SCALE : e->layers[li - 1].out_scale;
|
| 220 |
+
}
|
| 221 |
+
|
| 222 |
+
layer_forward(L, r, e->scratch_f);
|
| 223 |
+
|
| 224 |
+
if (L->residual) {
|
| 225 |
+
for (int c = 0; c < L->out_c; c++)
|
| 226 |
+
e->scratch_f[c] += (float)e->block_in[c] * e->block_in_scale;
|
| 227 |
+
}
|
| 228 |
+
if (L->relu) {
|
| 229 |
+
for (int c = 0; c < L->out_c; c++)
|
| 230 |
+
if (e->scratch_f[c] < 0.f) e->scratch_f[c] = 0.f;
|
| 231 |
+
}
|
| 232 |
+
if (li == PWW_NUM_LAYERS - 1) {
|
| 233 |
+
// final logits: snap to grid to mirror the simulator, then
|
| 234 |
+
// return as float
|
| 235 |
+
float s = L->out_scale;
|
| 236 |
+
for (int c = 0; c < L->out_c; c++) {
|
| 237 |
+
int8_t q = quant_clamp(e->scratch_f[c], 1.0f / s);
|
| 238 |
+
logits_out[c] = (float)q * s;
|
| 239 |
+
}
|
| 240 |
+
return 0;
|
| 241 |
+
}
|
| 242 |
+
float inv = 1.0f / L->out_scale;
|
| 243 |
+
int8_t *nxt = (cur == e->col_a) ? e->col_b : e->col_a;
|
| 244 |
+
for (int c = 0; c < L->out_c; c++)
|
| 245 |
+
nxt[c] = quant_clamp(e->scratch_f[c], inv);
|
| 246 |
+
if (pww_dump_layer == li) {
|
| 247 |
+
for (int c = 0; c < L->out_c && c < 256; c++)
|
| 248 |
+
logits_out[c] = (float)nxt[c];
|
| 249 |
+
return 0;
|
| 250 |
+
}
|
| 251 |
+
cur = nxt;
|
| 252 |
+
}
|
| 253 |
+
return -1; // unreachable
|
| 254 |
+
}
|
engine_c/pww_engine.h
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Streaming int8 phoneme engine - portable C (PC test build + ESP32-S3).
|
| 2 |
+
// Consumes one 40 ms feature frame at a time (40 log-mel values at the
|
| 3 |
+
// model's 20 ms output rate this is stride-2, so the caller feeds TWO
|
| 4 |
+
// 10 ms-hop mel frames per step); emits one logits vector per step.
|
| 5 |
+
//
|
| 6 |
+
// Design contract = phoneme_engine/quantize.py fake_quant_forward():
|
| 7 |
+
// - weights int8 per-output-channel, BN folded
|
| 8 |
+
// - int8 x int8 -> int32 accumulate, float requant (combined scale),
|
| 9 |
+
// residual added in float, ReLU, snap to int8 grid of the next layer
|
| 10 |
+
// - strictly causal: each layer keeps a ring buffer of its int8 input
|
| 11 |
+
// history, so per-step cost is one new column per layer
|
| 12 |
+
#pragma once
|
| 13 |
+
#include <stdint.h>
|
| 14 |
+
|
| 15 |
+
#ifdef __cplusplus
|
| 16 |
+
extern "C" {
|
| 17 |
+
#endif
|
| 18 |
+
|
| 19 |
+
#define PWW_MELS 40
|
| 20 |
+
|
| 21 |
+
typedef struct pww_engine pww_engine_t;
|
| 22 |
+
|
| 23 |
+
// Allocates all layer state (ring buffers) on the heap. Returns NULL on
|
| 24 |
+
// allocation failure.
|
| 25 |
+
pww_engine_t *pww_engine_create(void);
|
| 26 |
+
void pww_engine_destroy(pww_engine_t *e);
|
| 27 |
+
void pww_engine_reset(pww_engine_t *e);
|
| 28 |
+
|
| 29 |
+
// Feed TWO consecutive 10 ms-hop mel frames (each PWW_MELS floats, already
|
| 30 |
+
// EMA-normalized like features.py). Writes PWW_NUM_CLASSES float logits
|
| 31 |
+
// (log-prob differences are what the decoder consumes; absolute offset is
|
| 32 |
+
// meaningless). Returns 0 on success.
|
| 33 |
+
int pww_engine_step(pww_engine_t *e, const float *mel0, const float *mel1,
|
| 34 |
+
float *logits_out);
|
| 35 |
+
|
| 36 |
+
#ifdef __cplusplus
|
| 37 |
+
}
|
| 38 |
+
#endif
|
engine_c/pww_frontend.c
ADDED
|
@@ -0,0 +1,127 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Streaming log-mel frontend - C port of phoneme_engine/features.py.
|
| 2 |
+
// Frame t is centered at sample t*HOP (torchaudio center=True), so the
|
| 3 |
+
// stream carries WIN/2 = 200 samples (12.5 ms) of lookahead latency.
|
| 4 |
+
#include "pww_frontend.h"
|
| 5 |
+
|
| 6 |
+
#include <math.h>
|
| 7 |
+
#include <stdint.h>
|
| 8 |
+
#include <string.h>
|
| 9 |
+
|
| 10 |
+
#include "frontend_data.h"
|
| 11 |
+
|
| 12 |
+
#ifndef M_PI
|
| 13 |
+
#define M_PI 3.14159265358979323846
|
| 14 |
+
#endif
|
| 15 |
+
|
| 16 |
+
// ---- 512-point iterative radix-2 complex FFT (input real) -------------
|
| 17 |
+
#define NFFT PWW_FE_NFFT
|
| 18 |
+
|
| 19 |
+
static float fft_re[NFFT], fft_im[NFFT];
|
| 20 |
+
static float tw_cos[NFFT / 2], tw_sin[NFFT / 2];
|
| 21 |
+
static int fft_init_done = 0;
|
| 22 |
+
|
| 23 |
+
static void fft_init(void) {
|
| 24 |
+
for (int i = 0; i < NFFT / 2; i++) {
|
| 25 |
+
tw_cos[i] = cosf(-2.0f * (float)M_PI * i / NFFT);
|
| 26 |
+
tw_sin[i] = sinf(-2.0f * (float)M_PI * i / NFFT);
|
| 27 |
+
}
|
| 28 |
+
fft_init_done = 1;
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
static void fft512(void) {
|
| 32 |
+
// bit-reversal permutation
|
| 33 |
+
for (int i = 1, j = 0; i < NFFT; i++) {
|
| 34 |
+
int bit = NFFT >> 1;
|
| 35 |
+
for (; j & bit; bit >>= 1) j ^= bit;
|
| 36 |
+
j ^= bit;
|
| 37 |
+
if (i < j) {
|
| 38 |
+
float tr = fft_re[i]; fft_re[i] = fft_re[j]; fft_re[j] = tr;
|
| 39 |
+
float ti = fft_im[i]; fft_im[i] = fft_im[j]; fft_im[j] = ti;
|
| 40 |
+
}
|
| 41 |
+
}
|
| 42 |
+
for (int len = 2; len <= NFFT; len <<= 1) {
|
| 43 |
+
int half = len >> 1, step = NFFT / len;
|
| 44 |
+
for (int i = 0; i < NFFT; i += len) {
|
| 45 |
+
for (int k = 0; k < half; k++) {
|
| 46 |
+
float wr = tw_cos[k * step], wi = tw_sin[k * step];
|
| 47 |
+
int a = i + k, b = i + k + half;
|
| 48 |
+
float xr = fft_re[b] * wr - fft_im[b] * wi;
|
| 49 |
+
float xi = fft_re[b] * wi + fft_im[b] * wr;
|
| 50 |
+
fft_re[b] = fft_re[a] - xr;
|
| 51 |
+
fft_im[b] = fft_im[a] - xi;
|
| 52 |
+
fft_re[a] += xr;
|
| 53 |
+
fft_im[a] += xi;
|
| 54 |
+
}
|
| 55 |
+
}
|
| 56 |
+
}
|
| 57 |
+
}
|
| 58 |
+
|
| 59 |
+
// ---- streaming state ---------------------------------------------------
|
| 60 |
+
// ring of raw samples; enough for one centered window plus slack
|
| 61 |
+
#define SBUF (PWW_FE_WIN + 4 * PWW_FE_HOP)
|
| 62 |
+
|
| 63 |
+
struct pww_frontend {
|
| 64 |
+
float samples[SBUF];
|
| 65 |
+
int n_samples; // total pushed
|
| 66 |
+
int64_t next_frame; // next frame index to emit
|
| 67 |
+
float ema[PWW_FE_NMELS];
|
| 68 |
+
int ema_init;
|
| 69 |
+
};
|
| 70 |
+
|
| 71 |
+
static struct pww_frontend g_fe;
|
| 72 |
+
|
| 73 |
+
void pww_frontend_reset(void) {
|
| 74 |
+
memset(&g_fe, 0, sizeof(g_fe));
|
| 75 |
+
if (!fft_init_done) fft_init();
|
| 76 |
+
}
|
| 77 |
+
|
| 78 |
+
// Push samples; calls emit(mel_frame) for every completed 10 ms frame.
|
| 79 |
+
void pww_frontend_push(const float *x, int n,
|
| 80 |
+
void (*emit)(const float *mel, void *user),
|
| 81 |
+
void *user) {
|
| 82 |
+
struct pww_frontend *fe = &g_fe;
|
| 83 |
+
for (int i = 0; i < n; i++) {
|
| 84 |
+
fe->samples[fe->n_samples % SBUF] = x[i];
|
| 85 |
+
fe->n_samples++;
|
| 86 |
+
// frame f is ready when its centered window [f*HOP-200, f*HOP+200)
|
| 87 |
+
// is fully available
|
| 88 |
+
long long center = fe->next_frame * PWW_FE_HOP;
|
| 89 |
+
while (center + PWW_FE_WIN / 2 <= fe->n_samples) {
|
| 90 |
+
float mel[PWW_FE_NMELS];
|
| 91 |
+
long long start = center - PWW_FE_WIN / 2;
|
| 92 |
+
memset(fft_re, 0, sizeof(fft_re));
|
| 93 |
+
memset(fft_im, 0, sizeof(fft_im));
|
| 94 |
+
int off = (NFFT - PWW_FE_WIN) / 2; // window centered in FFT
|
| 95 |
+
for (int w = 0; w < PWW_FE_WIN; w++) {
|
| 96 |
+
long long s = start + w;
|
| 97 |
+
// reflect padding at the stream start (center=True)
|
| 98 |
+
if (s < 0) s = -s;
|
| 99 |
+
float v = (s < fe->n_samples)
|
| 100 |
+
? fe->samples[s % SBUF] : 0.f;
|
| 101 |
+
fft_re[off + w] = v * PWW_FE_HANN[w];
|
| 102 |
+
}
|
| 103 |
+
fft512();
|
| 104 |
+
for (int m = 0; m < PWW_FE_NMELS; m++) mel[m] = 0.f;
|
| 105 |
+
for (int b = 0; b < PWW_FE_NFREQS; b++) {
|
| 106 |
+
float p = fft_re[b] * fft_re[b] + fft_im[b] * fft_im[b];
|
| 107 |
+
const float *fbrow = PWW_FE_MELFB + b;
|
| 108 |
+
for (int m = 0; m < PWW_FE_NMELS; m++)
|
| 109 |
+
mel[m] += p * fbrow[(size_t)m * PWW_FE_NFREQS];
|
| 110 |
+
}
|
| 111 |
+
for (int m = 0; m < PWW_FE_NMELS; m++) {
|
| 112 |
+
float lm = logf(mel[m] + 1e-6f);
|
| 113 |
+
// causal EMA mean subtraction (features.py lfilter)
|
| 114 |
+
float e = fe->ema_init
|
| 115 |
+
? (1.f - PWW_FE_EMA_ALPHA) * fe->ema[m]
|
| 116 |
+
+ PWW_FE_EMA_ALPHA * lm
|
| 117 |
+
: PWW_FE_EMA_ALPHA * lm;
|
| 118 |
+
fe->ema[m] = e;
|
| 119 |
+
mel[m] = lm - e;
|
| 120 |
+
}
|
| 121 |
+
fe->ema_init = 1;
|
| 122 |
+
emit(mel, user);
|
| 123 |
+
fe->next_frame++;
|
| 124 |
+
center = fe->next_frame * PWW_FE_HOP;
|
| 125 |
+
}
|
| 126 |
+
}
|
| 127 |
+
}
|
engine_c/pww_frontend.h
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
// Streaming log-mel frontend matching the training features exactly.
|
| 2 |
+
#pragma once
|
| 3 |
+
|
| 4 |
+
#ifdef __cplusplus
|
| 5 |
+
extern "C" {
|
| 6 |
+
#endif
|
| 7 |
+
|
| 8 |
+
void pww_frontend_reset(void);
|
| 9 |
+
void pww_frontend_push(const float *x, int n,
|
| 10 |
+
void (*emit)(const float *mel, void *user),
|
| 11 |
+
void *user);
|
| 12 |
+
|
| 13 |
+
#ifdef __cplusplus
|
| 14 |
+
}
|
| 15 |
+
#endif
|
examples/neptuno_universal.json
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"phrase": "neptuno",
|
| 3 |
+
"spotters": [
|
| 4 |
+
{
|
| 5 |
+
"phones": [
|
| 6 |
+
"N",
|
| 7 |
+
"EH",
|
| 8 |
+
"P",
|
| 9 |
+
"T",
|
| 10 |
+
"UW",
|
| 11 |
+
"N",
|
| 12 |
+
"OW"
|
| 13 |
+
],
|
| 14 |
+
"source": "dictionary",
|
| 15 |
+
"threshold": -2.5
|
| 16 |
+
},
|
| 17 |
+
{
|
| 18 |
+
"phones": [
|
| 19 |
+
"N",
|
| 20 |
+
"EH",
|
| 21 |
+
"T",
|
| 22 |
+
"UW",
|
| 23 |
+
"N",
|
| 24 |
+
"OW"
|
| 25 |
+
],
|
| 26 |
+
"source": "population",
|
| 27 |
+
"threshold": -2.5
|
| 28 |
+
},
|
| 29 |
+
{
|
| 30 |
+
"phones": [
|
| 31 |
+
"N",
|
| 32 |
+
"AE",
|
| 33 |
+
"T",
|
| 34 |
+
"UW",
|
| 35 |
+
"N",
|
| 36 |
+
"OW"
|
| 37 |
+
],
|
| 38 |
+
"source": "population",
|
| 39 |
+
"threshold": -2.5
|
| 40 |
+
},
|
| 41 |
+
{
|
| 42 |
+
"phones": [
|
| 43 |
+
"N",
|
| 44 |
+
"EH",
|
| 45 |
+
"T",
|
| 46 |
+
"T",
|
| 47 |
+
"UW",
|
| 48 |
+
"N",
|
| 49 |
+
"OW"
|
| 50 |
+
],
|
| 51 |
+
"source": "population",
|
| 52 |
+
"threshold": -2.5
|
| 53 |
+
},
|
| 54 |
+
{
|
| 55 |
+
"phones": [
|
| 56 |
+
"N",
|
| 57 |
+
"EH",
|
| 58 |
+
"P",
|
| 59 |
+
"T",
|
| 60 |
+
"UW",
|
| 61 |
+
"N",
|
| 62 |
+
"OW"
|
| 63 |
+
],
|
| 64 |
+
"source": "population",
|
| 65 |
+
"threshold": -2.5
|
| 66 |
+
}
|
| 67 |
+
],
|
| 68 |
+
"mode": "universal"
|
| 69 |
+
}
|
frontend_data.h
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
model_int8.h
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
phoneme_engine/__init__.py
ADDED
|
File without changes
|
phoneme_engine/decoder.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Streaming keyword spotting over CTC phoneme posteriors.
|
| 2 |
+
|
| 3 |
+
Keyword-filler decoding with two robustness measures beyond the textbook
|
| 4 |
+
version (both port directly to C on the ESP32):
|
| 5 |
+
|
| 6 |
+
- minimum phone duration: every phone is expanded to two chained states,
|
| 7 |
+
so an alignment must spend >= 2 frames (40 ms) per phone. Kills
|
| 8 |
+
spurious single-frame matches.
|
| 9 |
+
- duration-normalized scoring: each Viterbi token carries the frame at
|
| 10 |
+
which its path entered the keyword, and the detection statistic is
|
| 11 |
+
(keyword_score - filler_score) / path_duration -- an average
|
| 12 |
+
per-frame deficit, comparable across phrase lengths and speaking
|
| 13 |
+
rates. Alignments faster than 2 frames/phone or slower than
|
| 14 |
+
MAX_FRAMES_PER_PHONE are rejected outright.
|
| 15 |
+
|
| 16 |
+
Enrollment of a new wake word is just: phones = text_to_phones(phrase).
|
| 17 |
+
State chain per phone i: A_i -> B_i (same label), optional blank after.
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
import numpy as np
|
| 21 |
+
|
| 22 |
+
from .phones import BLANK, PHONE_TO_ID, text_to_phones
|
| 23 |
+
|
| 24 |
+
NEG_INF = -1e30
|
| 25 |
+
# 200 ms per phone upper bound: fluent speech runs 40-120 ms/phone, and
|
| 26 |
+
# looser caps let garbage alignments crawl across continuous speech
|
| 27 |
+
# (observed 1.5-2.3 s "matches" on background TV at 400 ms/phone)
|
| 28 |
+
MAX_FRAMES_PER_PHONE = 10
|
| 29 |
+
|
| 30 |
+
# Highly confusable phone pairs: an alignment may match either member.
|
| 31 |
+
# Text pronunciations (CMUdict/G2P) use canonical vowels, but real and TTS
|
| 32 |
+
# speech often reduces them -- e.g. "orbit" is listed as AO R B AH T yet
|
| 33 |
+
# actually said as AO R B IH T. Without this, a sharper acoustic model
|
| 34 |
+
# *punishes* the mismatch harder.
|
| 35 |
+
CONFUSABLE = {
|
| 36 |
+
"AH": ("AH", "IH", "ER"), # schwa reduces/r-colors: orbit -> orbERt
|
| 37 |
+
"IH": ("IH", "AH"),
|
| 38 |
+
"ER": ("ER", "AH"),
|
| 39 |
+
"AO": ("AO", "AA"), # cot-caught merger and accent variation
|
| 40 |
+
"AE": ("AE", "AA"), # trap-father variation (sakura, tanaka)
|
| 41 |
+
"UH": ("UH", "UW"), # lax/tense u (book/boot neighbors)
|
| 42 |
+
"K": ("K", "G"), # stops voice between vowels: nakuma->naguma
|
| 43 |
+
"G": ("G", "K"),
|
| 44 |
+
"T": ("T", "D"), # also covers tapped/rolled R heard as D
|
| 45 |
+
"D": ("D", "T"),
|
| 46 |
+
}
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class KeywordSpotter:
|
| 50 |
+
def __init__(self, phrase, threshold=-4.4, refractory_frames=50,
|
| 51 |
+
strong_margin=1.5, strong_ratio=0.5, phones=None,
|
| 52 |
+
wildcards=None, wildcard_penalty=2.2):
|
| 53 |
+
"""threshold: average per-frame deficit tolerated along the keyword
|
| 54 |
+
path (log-prob units, <= 0; closer to 0 = stricter).
|
| 55 |
+
refractory_frames: minimum output frames (20 ms) between fires.
|
| 56 |
+
strong_margin/strong_ratio: evidence requirement -- at least
|
| 57 |
+
strong_ratio of the path's phone frames must have their phone
|
| 58 |
+
within strong_margin log-prob of the frame's best class. Prevents
|
| 59 |
+
noisy/mushy audio (where nothing is clearly heard) from firing on
|
| 60 |
+
average-score alone.
|
| 61 |
+
phones: explicit phoneme sequence (e.g. from voice enrollment);
|
| 62 |
+
overrides the dictionary pronunciation of `phrase`."""
|
| 63 |
+
if phones is None:
|
| 64 |
+
phones = [p for p in text_to_phones(phrase) if p in PHONE_TO_ID]
|
| 65 |
+
else:
|
| 66 |
+
phones = [p for p in phones if p in PHONE_TO_ID]
|
| 67 |
+
if len(phones) < 2:
|
| 68 |
+
raise ValueError(f"phrase too short to spot: {phrase!r}")
|
| 69 |
+
self.phrase = phrase
|
| 70 |
+
self.phones = phones
|
| 71 |
+
self.threshold = threshold
|
| 72 |
+
self.refractory_frames = refractory_frames
|
| 73 |
+
self.strong_margin = strong_margin
|
| 74 |
+
self.strong_ratio = strong_ratio
|
| 75 |
+
# Mismatch tolerance (default OFF): measured on universal clips
|
| 76 |
+
# vs negatives, wildcarding lifted impostor scores as much as
|
| 77 |
+
# genuine ones and REDUCED recall at zero-FA operating points
|
| 78 |
+
# (vucano negmax -4.14 -> -2.89). Kept only for future
|
| 79 |
+
# per-count-lattice experiments.
|
| 80 |
+
if wildcards is None:
|
| 81 |
+
wildcards = 0
|
| 82 |
+
self.max_wild = wildcards * 2 # budget in frames
|
| 83 |
+
self.wild_pen = wildcard_penalty
|
| 84 |
+
|
| 85 |
+
# Build states: leading blank, then per phone A,B (+ trailing blank)
|
| 86 |
+
labels = [BLANK]
|
| 87 |
+
self.preds = [[0]] # predecessor state ids (self-loop
|
| 88 |
+
entry = [True] # implied for every state)
|
| 89 |
+
for i, p in enumerate(phones):
|
| 90 |
+
pid = PHONE_TO_ID[p]
|
| 91 |
+
a = len(labels)
|
| 92 |
+
if i == 0:
|
| 93 |
+
labels.append(pid); self.preds.append([0]); entry.append(True)
|
| 94 |
+
else:
|
| 95 |
+
pre = [a - 1] # blank between phones
|
| 96 |
+
if labels[a - 2] != pid:
|
| 97 |
+
pre.append(a - 2) # skip blank (different phones only)
|
| 98 |
+
labels.append(pid); self.preds.append(pre); entry.append(False)
|
| 99 |
+
labels.append(pid) # B state: only from A
|
| 100 |
+
self.preds.append([a]); entry.append(False)
|
| 101 |
+
labels.append(BLANK) # blank after phone
|
| 102 |
+
self.preds.append([a + 1]); entry.append(False)
|
| 103 |
+
self.labels = labels
|
| 104 |
+
# allowed emission ids per state (confusable vowels match either)
|
| 105 |
+
self.allowed = []
|
| 106 |
+
id_to_phone = {v: k for k, v in PHONE_TO_ID.items()}
|
| 107 |
+
for lab in labels:
|
| 108 |
+
if lab == BLANK:
|
| 109 |
+
self.allowed.append((BLANK,))
|
| 110 |
+
else:
|
| 111 |
+
ph = id_to_phone[lab]
|
| 112 |
+
self.allowed.append(tuple(
|
| 113 |
+
PHONE_TO_ID[p] for p in CONFUSABLE.get(ph, (ph,))))
|
| 114 |
+
self.entry = entry
|
| 115 |
+
self.n_states = len(labels)
|
| 116 |
+
self.finals = [self.n_states - 1, self.n_states - 2]
|
| 117 |
+
# min: ~60 ms per phone on average (individual phones may be
|
| 118 |
+
# shorter); max: 400 ms per phone. Anything outside is not a
|
| 119 |
+
# human saying the phrase.
|
| 120 |
+
self.min_dur = 3 * len(phones)
|
| 121 |
+
self.max_dur = MAX_FRAMES_PER_PHONE * len(phones)
|
| 122 |
+
self.reset()
|
| 123 |
+
|
| 124 |
+
def reset(self):
|
| 125 |
+
self.rel = np.full(self.n_states, NEG_INF)
|
| 126 |
+
self.start = np.zeros(self.n_states, dtype=np.int64)
|
| 127 |
+
self.pframes = np.zeros(self.n_states, dtype=np.int64)
|
| 128 |
+
self.strong = np.zeros(self.n_states, dtype=np.int64)
|
| 129 |
+
self.wild = np.zeros(self.n_states, dtype=np.int64)
|
| 130 |
+
self.cooldown = 0
|
| 131 |
+
self.t = 0
|
| 132 |
+
|
| 133 |
+
def _update(self, log_probs_frame):
|
| 134 |
+
"""One Viterbi DP step. Returns the best duration-valid normalized
|
| 135 |
+
score at a final state this frame (or None). No side effects on
|
| 136 |
+
detection state."""
|
| 137 |
+
lp = log_probs_frame
|
| 138 |
+
filler = float(lp.max())
|
| 139 |
+
prev_rel, prev_start = self.rel, self.start
|
| 140 |
+
prev_pf, prev_strong = self.pframes, self.strong
|
| 141 |
+
prev_wild = self.wild
|
| 142 |
+
cur_rel = np.full(self.n_states, NEG_INF)
|
| 143 |
+
cur_start = np.zeros(self.n_states, dtype=np.int64)
|
| 144 |
+
cur_pf = np.zeros(self.n_states, dtype=np.int64)
|
| 145 |
+
cur_strong = np.zeros(self.n_states, dtype=np.int64)
|
| 146 |
+
cur_wild = np.zeros(self.n_states, dtype=np.int64)
|
| 147 |
+
for s in range(self.n_states):
|
| 148 |
+
best = prev_rel[s]
|
| 149 |
+
best_start, best_pf, best_sf, best_w = (
|
| 150 |
+
prev_start[s], prev_pf[s], prev_strong[s], prev_wild[s])
|
| 151 |
+
for q in self.preds[s]:
|
| 152 |
+
if prev_rel[q] > best:
|
| 153 |
+
best = prev_rel[q]
|
| 154 |
+
best_start, best_pf, best_sf, best_w = (
|
| 155 |
+
prev_start[q], prev_pf[q], prev_strong[q],
|
| 156 |
+
prev_wild[q])
|
| 157 |
+
if self.entry[s] and 0.0 >= best:
|
| 158 |
+
# (re)start the keyword here; >= keeps the start time fresh
|
| 159 |
+
# while idling in silence, so duration stays meaningful
|
| 160 |
+
best, best_start, best_pf, best_sf, best_w = \
|
| 161 |
+
0.0, self.t, 0, 0, 0
|
| 162 |
+
emit = max(float(lp[i]) for i in self.allowed[s])
|
| 163 |
+
is_phone = self.labels[s] != BLANK
|
| 164 |
+
used_wild = 0
|
| 165 |
+
if is_phone and best_w < self.max_wild:
|
| 166 |
+
# mismatch tolerance: accept the frame's best class at a
|
| 167 |
+
# penalty when the expected phone isn't there (bounded
|
| 168 |
+
# budget converts exact matching into similarity matching)
|
| 169 |
+
soft = filler - self.wild_pen
|
| 170 |
+
if soft > emit:
|
| 171 |
+
emit = soft
|
| 172 |
+
used_wild = 1
|
| 173 |
+
cur_rel[s] = best + emit - filler
|
| 174 |
+
cur_start[s] = best_start
|
| 175 |
+
cur_pf[s] = best_pf + (1 if is_phone else 0)
|
| 176 |
+
cur_strong[s] = best_sf + (
|
| 177 |
+
1 if is_phone and emit >= filler - self.strong_margin else 0)
|
| 178 |
+
cur_wild[s] = best_w + used_wild
|
| 179 |
+
self.rel, self.start = cur_rel, cur_start
|
| 180 |
+
self.pframes, self.strong = cur_pf, cur_strong
|
| 181 |
+
self.wild = cur_wild
|
| 182 |
+
self.t += 1
|
| 183 |
+
norm_best = None
|
| 184 |
+
for s in self.finals:
|
| 185 |
+
dur = self.t - cur_start[s]
|
| 186 |
+
pf = cur_pf[s]
|
| 187 |
+
if not (self.min_dur <= dur <= self.max_dur):
|
| 188 |
+
continue
|
| 189 |
+
if pf < 2 * len(self.phones):
|
| 190 |
+
continue
|
| 191 |
+
# evidence requirement: the phrase's phones must actually have
|
| 192 |
+
# been the near-top hypothesis for enough of the match, not
|
| 193 |
+
# merely "not too costly on average" (mushy noisy audio)
|
| 194 |
+
if cur_strong[s] < self.strong_ratio * pf:
|
| 195 |
+
continue
|
| 196 |
+
# normalize by phone frames only: time spent in blank states is
|
| 197 |
+
# free in silence, so counting it would let blank-padded paths
|
| 198 |
+
# dilute their deficit below any threshold
|
| 199 |
+
norm = cur_rel[s] / pf
|
| 200 |
+
if norm_best is None or norm > norm_best[0]:
|
| 201 |
+
norm_best = (float(norm), int(dur))
|
| 202 |
+
return norm_best
|
| 203 |
+
|
| 204 |
+
def step(self, log_probs_frame):
|
| 205 |
+
"""Consume one frame of log-probs (C,). Returns (score, duration)
|
| 206 |
+
if the keyword fired this frame, else None. duration is in output
|
| 207 |
+
frames (20 ms each), for energy-gating by the caller."""
|
| 208 |
+
norm = self._update(log_probs_frame)
|
| 209 |
+
if self.cooldown > 0:
|
| 210 |
+
self.cooldown -= 1
|
| 211 |
+
return None
|
| 212 |
+
if norm is not None and norm[0] > self.threshold:
|
| 213 |
+
self.cooldown = self.refractory_frames
|
| 214 |
+
self.rel = np.full(self.n_states, NEG_INF)
|
| 215 |
+
return norm
|
| 216 |
+
return None
|
| 217 |
+
|
| 218 |
+
def run(self, log_probs):
|
| 219 |
+
"""Offline helper: (T, C) log-probs -> list of (frame, score)."""
|
| 220 |
+
hits = []
|
| 221 |
+
for t in range(log_probs.shape[0]):
|
| 222 |
+
s = self.step(log_probs[t])
|
| 223 |
+
if s is not None:
|
| 224 |
+
hits.append((t, s[0]))
|
| 225 |
+
return hits
|
| 226 |
+
|
| 227 |
+
def best_score(self, log_probs):
|
| 228 |
+
"""Max normalized final-state score over a clip (pure Viterbi, no
|
| 229 |
+
firing/reset side effects) -- used for calibration."""
|
| 230 |
+
self.reset()
|
| 231 |
+
best = -np.inf
|
| 232 |
+
for t in range(log_probs.shape[0]):
|
| 233 |
+
s = self._update(log_probs[t])
|
| 234 |
+
if s is not None and s[0] > best:
|
| 235 |
+
best = s[0]
|
| 236 |
+
return best
|
phoneme_engine/enroll.py
ADDED
|
@@ -0,0 +1,254 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Voice enrollment: learn YOUR pronunciation of a wake word in seconds.
|
| 2 |
+
|
| 3 |
+
Records the phrase three times, extracts the phoneme sequence the model
|
| 4 |
+
actually hears for your voice/accent/mic, and saves the variants alongside
|
| 5 |
+
the dictionary pronunciation:
|
| 6 |
+
|
| 7 |
+
python -m phoneme_engine.enroll "hey orbit"
|
| 8 |
+
|
| 9 |
+
Then use it live:
|
| 10 |
+
|
| 11 |
+
python -m phoneme_engine.live_demo "hey orbit" --enrolled
|
| 12 |
+
|
| 13 |
+
No training happens -- enrollment is just phoneme sequences (a few bytes),
|
| 14 |
+
same as typed enrollment, so it stays instant and fully on-device later.
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
import argparse
|
| 18 |
+
import json
|
| 19 |
+
import time
|
| 20 |
+
from pathlib import Path
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
import sounddevice as sd
|
| 24 |
+
import soundfile as sf
|
| 25 |
+
import torch
|
| 26 |
+
|
| 27 |
+
from .decoder import KeywordSpotter
|
| 28 |
+
from .features import SAMPLE_RATE, LogMel
|
| 29 |
+
from .model import PhonemeTCN
|
| 30 |
+
from .phones import BLANK, ids_to_phones
|
| 31 |
+
|
| 32 |
+
ROOT = Path(__file__).resolve().parent.parent
|
| 33 |
+
ENROLL_DIR = ROOT / "enrollments"
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def slug(phrase):
|
| 37 |
+
return "".join(c if c.isalnum() else "_" for c in phrase.lower())
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def min_variant_len(n_dict_phones):
|
| 41 |
+
"""Variants much shorter than the real phrase false-trigger constantly
|
| 42 |
+
(a 3-phone variant matches half of English)."""
|
| 43 |
+
return max(4, round(0.6 * n_dict_phones))
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def phone_segment(lp, merge_gap=15, margin=3):
|
| 47 |
+
"""(t0, t1) output-frame span of the densest contiguous run of frames
|
| 48 |
+
where the model hears an actual phoneme (blank prob < 0.5). Isolates
|
| 49 |
+
the spoken phrase from background sounds elsewhere in the take."""
|
| 50 |
+
p_blank = np.exp(lp[:, BLANK])
|
| 51 |
+
active = p_blank < 0.5
|
| 52 |
+
runs = []
|
| 53 |
+
cur = None
|
| 54 |
+
n = len(active)
|
| 55 |
+
for i in range(n + 1):
|
| 56 |
+
a = active[i] if i < n else False
|
| 57 |
+
if a and cur is None:
|
| 58 |
+
cur = i
|
| 59 |
+
elif not a and cur is not None:
|
| 60 |
+
if i + merge_gap <= n and active[i:i + merge_gap].any():
|
| 61 |
+
continue
|
| 62 |
+
runs.append((cur, i))
|
| 63 |
+
cur = None
|
| 64 |
+
if not runs:
|
| 65 |
+
return 0, lp.shape[0]
|
| 66 |
+
# densest = most active frames inside the run
|
| 67 |
+
t0, t1 = max(runs, key=lambda r: active[r[0]:r[1]].sum())
|
| 68 |
+
return max(0, t0 - margin), min(lp.shape[0], t1 + margin)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def extract_phones(lp, segment=True):
|
| 72 |
+
"""Greedy phoneme sequence from (T, C) log-probs, ignoring blanks.
|
| 73 |
+
With segment=True only the densest phoneme burst is transcribed
|
| 74 |
+
(background-proof)."""
|
| 75 |
+
t0, t1 = phone_segment(lp) if segment else (0, lp.shape[0])
|
| 76 |
+
ids = lp[t0:t1].argmax(axis=1).tolist()
|
| 77 |
+
seq = []
|
| 78 |
+
prev = None
|
| 79 |
+
for i in ids:
|
| 80 |
+
if i != prev and i != BLANK:
|
| 81 |
+
seq.append(i)
|
| 82 |
+
prev = i
|
| 83 |
+
return ids_to_phones(seq)
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def main():
|
| 87 |
+
ap = argparse.ArgumentParser()
|
| 88 |
+
ap.add_argument("phrase")
|
| 89 |
+
ap.add_argument("--takes", type=int, default=3)
|
| 90 |
+
ap.add_argument("--seconds", type=float, default=4.0)
|
| 91 |
+
ap.add_argument("--ckpt", default=str(ROOT / "checkpoints" / "best.pt"))
|
| 92 |
+
args = ap.parse_args()
|
| 93 |
+
|
| 94 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 95 |
+
frontend = LogMel().to(device).eval()
|
| 96 |
+
model = PhonemeTCN().to(device).eval()
|
| 97 |
+
state = torch.load(args.ckpt, map_location=device, weights_only=True)
|
| 98 |
+
model.load_state_dict(state["model"])
|
| 99 |
+
|
| 100 |
+
dict_spotter = KeywordSpotter(args.phrase)
|
| 101 |
+
print(f"dictionary pronunciation: {' '.join(dict_spotter.phones)}")
|
| 102 |
+
|
| 103 |
+
variants = []
|
| 104 |
+
take_lps = []
|
| 105 |
+
ENROLL_DIR.mkdir(exist_ok=True)
|
| 106 |
+
for take in range(1, args.takes + 1):
|
| 107 |
+
# capture starts before the countdown so the stream is already
|
| 108 |
+
# open when the user speaks (stream startup used to clip speech)
|
| 109 |
+
total = 3 * 0.7 + args.seconds
|
| 110 |
+
audio = sd.rec(int(total * SAMPLE_RATE),
|
| 111 |
+
samplerate=SAMPLE_RATE, channels=1, dtype="float32")
|
| 112 |
+
print(f"\ntake {take}/{args.takes}: say {args.phrase!r} once, "
|
| 113 |
+
"normally")
|
| 114 |
+
for i in (3, 2, 1):
|
| 115 |
+
print(f" {i}...", flush=True)
|
| 116 |
+
time.sleep(0.7)
|
| 117 |
+
print(" SPEAK NOW", flush=True)
|
| 118 |
+
sd.wait()
|
| 119 |
+
audio = audio[:, 0]
|
| 120 |
+
peak_db = 20 * np.log10(np.abs(audio).max() + 1e-12)
|
| 121 |
+
if peak_db < -30:
|
| 122 |
+
print(f" !! too quiet (peak {peak_db:.1f} dBFS) - retake")
|
| 123 |
+
continue
|
| 124 |
+
sf.write(str(ENROLL_DIR / f"{slug(args.phrase)}_take{take}.wav"),
|
| 125 |
+
audio, SAMPLE_RATE)
|
| 126 |
+
wav = torch.from_numpy(audio).unsqueeze(0).to(device)
|
| 127 |
+
with torch.no_grad():
|
| 128 |
+
lp = torch.log_softmax(model(frontend(wav)).float(),
|
| 129 |
+
dim=2)[0].cpu().numpy()
|
| 130 |
+
phones = extract_phones(lp)
|
| 131 |
+
print(f" heard: {' '.join(phones) or '(nothing)'}")
|
| 132 |
+
n_dict = len(dict_spotter.phones)
|
| 133 |
+
if len(phones) < min_variant_len(n_dict):
|
| 134 |
+
print(" !! too few phonemes heard (short variants false-"
|
| 135 |
+
"trigger constantly) - retake ignored")
|
| 136 |
+
continue
|
| 137 |
+
if len(phones) > 2 * n_dict + 4:
|
| 138 |
+
print(" !! heard too much (background speech?) - retake ignored")
|
| 139 |
+
continue
|
| 140 |
+
variants.append(phones)
|
| 141 |
+
take_lps.append(lp)
|
| 142 |
+
|
| 143 |
+
if not variants:
|
| 144 |
+
print("\nno usable takes - nothing saved")
|
| 145 |
+
return
|
| 146 |
+
|
| 147 |
+
# dedupe identical variants
|
| 148 |
+
uniq = []
|
| 149 |
+
for v in variants:
|
| 150 |
+
if v not in uniq:
|
| 151 |
+
uniq.append(v)
|
| 152 |
+
|
| 153 |
+
spotters = calibrate_spotters(args.phrase, uniq, take_lps, model,
|
| 154 |
+
frontend, device)
|
| 155 |
+
out = {
|
| 156 |
+
"phrase": args.phrase,
|
| 157 |
+
"spotters": spotters,
|
| 158 |
+
}
|
| 159 |
+
path = ENROLL_DIR / f"{slug(args.phrase)}.json"
|
| 160 |
+
path.write_text(json.dumps(out, indent=2), encoding="utf-8")
|
| 161 |
+
print(f"\nsaved {len(spotters)} calibrated spotter(s) -> {path}")
|
| 162 |
+
print("test live with:\n python -m phoneme_engine.live_demo "
|
| 163 |
+
f"\"{args.phrase}\" --enrolled")
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def calibrate_spotters(phrase, variants, take_lps, model, frontend, device,
|
| 167 |
+
n_negatives=60):
|
| 168 |
+
"""Per-variant thresholds: each pronunciation (dictionary + voice
|
| 169 |
+
variants) is calibrated independently against the user's takes and a
|
| 170 |
+
sample of negative speech. Variants that cannot separate the user's
|
| 171 |
+
voice from negatives are dropped -- a loose variant with a shared
|
| 172 |
+
threshold poisons the whole enrollment."""
|
| 173 |
+
candidates = [("dictionary", None)] + [("voice", v) for v in variants]
|
| 174 |
+
|
| 175 |
+
neg_lps = []
|
| 176 |
+
neg_dir = ROOT / "data" / "LibriSpeech" / "dev-clean"
|
| 177 |
+
if neg_dir.exists():
|
| 178 |
+
from .spot_file import load_wav
|
| 179 |
+
files = sorted(neg_dir.glob("**/*.flac"))
|
| 180 |
+
rng = np.random.default_rng(0)
|
| 181 |
+
idx = rng.choice(len(files), min(n_negatives, len(files)),
|
| 182 |
+
replace=False)
|
| 183 |
+
with torch.no_grad():
|
| 184 |
+
for i in sorted(idx):
|
| 185 |
+
wav = load_wav(files[i], device)
|
| 186 |
+
neg_lps.append(torch.log_softmax(
|
| 187 |
+
model(frontend(wav)).float(), dim=2)[0].cpu().numpy())
|
| 188 |
+
|
| 189 |
+
out = []
|
| 190 |
+
fallback = None
|
| 191 |
+
print()
|
| 192 |
+
for source, phones in candidates:
|
| 193 |
+
sp = KeywordSpotter(phrase, phones=phones)
|
| 194 |
+
scores = [sp.best_score(lp) for lp in take_lps]
|
| 195 |
+
pos_best, pos_worst = max(scores), min(scores)
|
| 196 |
+
neg = max((sp.best_score(lp) for lp in neg_lps), default=-np.inf)
|
| 197 |
+
label = f"{source}: {' '.join(sp.phones)}"
|
| 198 |
+
if pos_best - neg < 0.75:
|
| 199 |
+
print(f" dropped {label} (your best {pos_best:.2f} vs "
|
| 200 |
+
f"negative {neg:.2f} - too false-alarm prone)")
|
| 201 |
+
if pos_best != -np.inf:
|
| 202 |
+
strict = round(float(min(-2.5, max(-6.0, neg + 0.3))), 2)
|
| 203 |
+
cand = {"phones": list(sp.phones), "source": source,
|
| 204 |
+
"threshold": strict}
|
| 205 |
+
if fallback is None or pos_best - neg > fallback[0]:
|
| 206 |
+
fallback = (pos_best - neg, cand)
|
| 207 |
+
continue
|
| 208 |
+
# anchor on the WEAKEST take (live speech varies more than
|
| 209 |
+
# enrollment takes), floored above the hardest negative
|
| 210 |
+
thr = max(pos_worst - 0.5, neg + 0.3) if neg_lps else pos_worst - 0.5
|
| 211 |
+
thr = round(float(min(-2.5, max(-6.0, thr))), 2)
|
| 212 |
+
print(f" kept {label} threshold {thr:.2f} "
|
| 213 |
+
f"(your takes {pos_worst:.2f}..{pos_best:.2f}, "
|
| 214 |
+
f"hardest negative {neg:.2f})")
|
| 215 |
+
entry = {"phones": list(sp.phones), "source": source,
|
| 216 |
+
"threshold": thr}
|
| 217 |
+
out.append(entry)
|
| 218 |
+
if fallback is None or pos_best - neg > fallback[0]:
|
| 219 |
+
fallback = (pos_best - neg, entry)
|
| 220 |
+
|
| 221 |
+
if not out and fallback is not None:
|
| 222 |
+
print(" (all candidates marginal - keeping the best one strictly)")
|
| 223 |
+
out.append(fallback[1])
|
| 224 |
+
return out
|
| 225 |
+
|
| 226 |
+
|
| 227 |
+
def load_spotter_configs(phrase, default_threshold=-4.4, profile=None):
|
| 228 |
+
"""Returns [(phones, threshold), ...] for a phrase: the calibrated
|
| 229 |
+
enrollment if present, else just the dictionary pronunciation.
|
| 230 |
+
profile selects an alternate enrollment file, e.g. 'universal'."""
|
| 231 |
+
name = slug(phrase) + (f"_{profile}" if profile else "")
|
| 232 |
+
path = ENROLL_DIR / f"{name}.json"
|
| 233 |
+
dict_phones = list(KeywordSpotter(phrase).phones)
|
| 234 |
+
if path.exists():
|
| 235 |
+
data = json.loads(path.read_text(encoding="utf-8"))
|
| 236 |
+
if "spotters" in data:
|
| 237 |
+
configs = []
|
| 238 |
+
min_len = min_variant_len(len(dict_phones))
|
| 239 |
+
for s in data["spotters"]:
|
| 240 |
+
if s["source"] == "dictionary" or len(s["phones"]) >= min_len:
|
| 241 |
+
configs.append((s["phones"], s["threshold"]))
|
| 242 |
+
if configs:
|
| 243 |
+
return configs
|
| 244 |
+
# legacy schema: variants without per-variant thresholds
|
| 245 |
+
configs = [(dict_phones, default_threshold)]
|
| 246 |
+
for v in data.get("variants", []):
|
| 247 |
+
if len(v) >= min_variant_len(len(dict_phones)):
|
| 248 |
+
configs.append((v, default_threshold))
|
| 249 |
+
return configs
|
| 250 |
+
return [(dict_phones, default_threshold)]
|
| 251 |
+
|
| 252 |
+
|
| 253 |
+
if __name__ == "__main__":
|
| 254 |
+
main()
|
phoneme_engine/enroll_universal.py
ADDED
|
@@ -0,0 +1,183 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Universal (speaker-independent) enrollment: build a wake word config
|
| 2 |
+
from the *population's* pronunciations instead of one person's.
|
| 3 |
+
|
| 4 |
+
Synthesizes the phrase with many diverse TTS voices, keeps only
|
| 5 |
+
Whisper-verified clips, extracts each clip's observed phoneme variant,
|
| 6 |
+
keeps variants that recur across different voices, and calibrates every
|
| 7 |
+
spotter against negative speech. Zero training; output is the same
|
| 8 |
+
~50-byte-per-spotter config schema as personal enrollment.
|
| 9 |
+
|
| 10 |
+
python -m phoneme_engine.enroll_universal "sakura" --say "sah koo rah"
|
| 11 |
+
|
| 12 |
+
Then: python -m phoneme_engine.live_demo "sakura" --enrolled --profile universal
|
| 13 |
+
"""
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import json
|
| 17 |
+
import subprocess
|
| 18 |
+
import sys
|
| 19 |
+
from collections import Counter
|
| 20 |
+
from pathlib import Path
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
import soundfile as sf
|
| 24 |
+
import torch
|
| 25 |
+
|
| 26 |
+
from .decoder import KeywordSpotter
|
| 27 |
+
from .enroll import ENROLL_DIR, extract_phones, min_variant_len, slug
|
| 28 |
+
from .features import LogMel
|
| 29 |
+
from .model import PhonemeTCN
|
| 30 |
+
|
| 31 |
+
ROOT = Path(__file__).resolve().parent.parent
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def _norm(s):
|
| 35 |
+
return "".join(c for c in s.lower() if c.isalpha())
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def _lev(a, b):
|
| 39 |
+
prev = list(range(len(b) + 1))
|
| 40 |
+
for i, x in enumerate(a, 1):
|
| 41 |
+
cur = [i]
|
| 42 |
+
for j, y in enumerate(b, 1):
|
| 43 |
+
cur.append(min(prev[j] + 1, cur[j - 1] + 1,
|
| 44 |
+
prev[j - 1] + (x != y)))
|
| 45 |
+
prev = cur
|
| 46 |
+
return prev[-1]
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
_DIGITS = {"0": "zero", "1": "one", "2": "two", "3": "three", "4": "four",
|
| 50 |
+
"5": "five", "6": "six", "7": "seven", "8": "eight",
|
| 51 |
+
"9": "nine"}
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def transcript_matches(text, phrase, say, target_phones):
|
| 55 |
+
"""Letter-level OR phoneme-level match. Whisper spells invented words
|
| 56 |
+
unpredictably ('Neptune-O', 'NEP 2 NO'), but its spellings are
|
| 57 |
+
usually phonetically faithful - so G2P the transcript and compare
|
| 58 |
+
phone sequences."""
|
| 59 |
+
from .phones import text_to_phones
|
| 60 |
+
t_letters = _norm(text)
|
| 61 |
+
for tgt in (_norm(phrase), _norm(say)):
|
| 62 |
+
if _lev(t_letters, tgt) <= max(1, len(tgt) // 4):
|
| 63 |
+
return True
|
| 64 |
+
spoken = "".join(_DIGITS.get(c, c) if c.isdigit() else c
|
| 65 |
+
for c in text.lower())
|
| 66 |
+
t_phones = text_to_phones(spoken)
|
| 67 |
+
tol = max(1, (len(target_phones) + 2) // 3)
|
| 68 |
+
return _lev(t_phones, list(target_phones)) <= tol
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def main():
|
| 72 |
+
ap = argparse.ArgumentParser()
|
| 73 |
+
ap.add_argument("phrase")
|
| 74 |
+
ap.add_argument("--say", default=None,
|
| 75 |
+
help="phonetic respelling for TTS (exotic words get "
|
| 76 |
+
"anglicized otherwise, e.g. sakura -> secure-ah)")
|
| 77 |
+
ap.add_argument("--n", type=int, default=36)
|
| 78 |
+
ap.add_argument("--min-votes", type=int, default=3,
|
| 79 |
+
help="a variant must appear in this many different "
|
| 80 |
+
"voices to count as population-level")
|
| 81 |
+
ap.add_argument("--ckpt", default=str(ROOT / "checkpoints" / "best.pt"))
|
| 82 |
+
args = ap.parse_args()
|
| 83 |
+
|
| 84 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 85 |
+
frontend = LogMel().to(device).eval()
|
| 86 |
+
model = PhonemeTCN().to(device).eval()
|
| 87 |
+
state = torch.load(args.ckpt, map_location=device, weights_only=True)
|
| 88 |
+
model.load_state_dict(state["model"])
|
| 89 |
+
|
| 90 |
+
# 1. synthesize across diverse voices
|
| 91 |
+
say = args.say or args.phrase
|
| 92 |
+
clips = ROOT / "enrollments" / f"_univ_{slug(args.phrase)}"
|
| 93 |
+
if not clips.exists() or len(list(clips.glob("*.wav"))) < args.n:
|
| 94 |
+
subprocess.run([sys.executable, "-m",
|
| 95 |
+
"phoneme_engine.make_piper_clips", say,
|
| 96 |
+
"--n", str(args.n), "--out", str(clips),
|
| 97 |
+
"--seed", "17"], check=True)
|
| 98 |
+
|
| 99 |
+
# 2. Whisper-verify
|
| 100 |
+
from transformers import pipeline
|
| 101 |
+
asr = pipeline("automatic-speech-recognition",
|
| 102 |
+
model="openai/whisper-small", device=0,
|
| 103 |
+
dtype=torch.float16)
|
| 104 |
+
from .phones import text_to_phones
|
| 105 |
+
target_phones = text_to_phones(say)
|
| 106 |
+
verified = []
|
| 107 |
+
for f in sorted(clips.glob("*.wav")):
|
| 108 |
+
audio, sr = sf.read(f, dtype="float32")
|
| 109 |
+
text = asr({"raw": audio, "sampling_rate": sr})["text"]
|
| 110 |
+
if transcript_matches(text, args.phrase, say, target_phones):
|
| 111 |
+
verified.append(f)
|
| 112 |
+
print(f"verified {len(verified)}/{args.n} clips")
|
| 113 |
+
if len(verified) < 10:
|
| 114 |
+
print("too few verified clips - check the respelling")
|
| 115 |
+
return
|
| 116 |
+
|
| 117 |
+
# 3. extract variants + vote across voices
|
| 118 |
+
lps = []
|
| 119 |
+
votes = Counter()
|
| 120 |
+
with torch.no_grad():
|
| 121 |
+
for f in verified:
|
| 122 |
+
audio, _ = sf.read(f, dtype="float32")
|
| 123 |
+
wav = torch.from_numpy(audio).float().unsqueeze(0).to(device)
|
| 124 |
+
lp = torch.log_softmax(model(frontend(wav)).float(),
|
| 125 |
+
dim=2)[0].cpu().numpy()
|
| 126 |
+
lps.append(lp)
|
| 127 |
+
v = tuple(extract_phones(lp))
|
| 128 |
+
n_dict = len(KeywordSpotter(args.phrase).phones)
|
| 129 |
+
if min_variant_len(n_dict) <= len(v) <= 2 * n_dict + 4:
|
| 130 |
+
votes[v] += 1
|
| 131 |
+
pop_variants = [list(v) for v, c in votes.most_common()
|
| 132 |
+
if c >= args.min_votes]
|
| 133 |
+
print("population variants:",
|
| 134 |
+
[(" ".join(v), votes[tuple(v)]) for v in pop_variants])
|
| 135 |
+
|
| 136 |
+
# 4. calibrate dictionary + population variants against negatives,
|
| 137 |
+
# with the verified clips as positives
|
| 138 |
+
z = np.load(ROOT / "bench_results" / "lp_neg_ls.npz",
|
| 139 |
+
allow_pickle=True)
|
| 140 |
+
neg_lps = list(z["lps"])[:200]
|
| 141 |
+
|
| 142 |
+
spotters = []
|
| 143 |
+
for source, phones in ([("dictionary", None)]
|
| 144 |
+
+ [("population", v) for v in pop_variants]):
|
| 145 |
+
sp = KeywordSpotter(args.phrase, phones=phones)
|
| 146 |
+
pos = sorted(sp.best_score(lp) for lp in lps)
|
| 147 |
+
neg = max(sp.best_score(lp) for lp in neg_lps)
|
| 148 |
+
pos_med = pos[len(pos) // 2]
|
| 149 |
+
if pos_med - neg < 0.5:
|
| 150 |
+
print(f" dropped {source}: {' '.join(sp.phones)} "
|
| 151 |
+
f"(median {pos_med:.2f} vs neg {neg:.2f})")
|
| 152 |
+
continue
|
| 153 |
+
# universal threshold: catch the median speaker, stay above the
|
| 154 |
+
# hardest negative
|
| 155 |
+
thr = float(min(-2.5, max(-6.0, max(pos_med - 0.4, neg + 0.3))))
|
| 156 |
+
kept = sum(1 for s in pos if s > thr) / len(pos)
|
| 157 |
+
print(f" kept {source}: {' '.join(sp.phones)} thr {thr:.2f} "
|
| 158 |
+
f"(covers {kept:.0%} of voices; hardest neg {neg:.2f})")
|
| 159 |
+
spotters.append({"phones": list(sp.phones), "source": source,
|
| 160 |
+
"threshold": thr})
|
| 161 |
+
|
| 162 |
+
if not spotters:
|
| 163 |
+
print("phrase refused: no spotter separable from negatives")
|
| 164 |
+
return
|
| 165 |
+
|
| 166 |
+
# coverage of the union
|
| 167 |
+
def hit(lp):
|
| 168 |
+
return any(KeywordSpotter(args.phrase, phones=s["phones"],
|
| 169 |
+
threshold=s["threshold"])
|
| 170 |
+
.best_score(lp) > s["threshold"] for s in spotters)
|
| 171 |
+
union = sum(1 for lp in lps if hit(lp)) / len(lps)
|
| 172 |
+
print(f"union coverage across {len(lps)} voices: {union:.0%}")
|
| 173 |
+
|
| 174 |
+
out = {"phrase": args.phrase, "spotters": spotters,
|
| 175 |
+
"mode": "universal"}
|
| 176 |
+
ENROLL_DIR.mkdir(exist_ok=True)
|
| 177 |
+
path = ENROLL_DIR / f"{slug(args.phrase)}_universal.json"
|
| 178 |
+
path.write_text(json.dumps(out, indent=2), encoding="utf-8")
|
| 179 |
+
print(f"saved -> {path}")
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
if __name__ == "__main__":
|
| 183 |
+
main()
|
phoneme_engine/features.py
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Log-mel feature frontend.
|
| 2 |
+
|
| 3 |
+
40 log-mel bands, 16 kHz, 25 ms window / 10 ms hop. Deliberately matches
|
| 4 |
+
what is cheap to compute on the ESP32-S3 later (esp-dsp / TFLite-Micro
|
| 5 |
+
audio frontend), so the on-device features can be made to line up with
|
| 6 |
+
training.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
import torch
|
| 10 |
+
import torchaudio
|
| 11 |
+
|
| 12 |
+
SAMPLE_RATE = 16000
|
| 13 |
+
N_MELS = 40
|
| 14 |
+
WIN_LENGTH = 400 # 25 ms
|
| 15 |
+
HOP_LENGTH = 160 # 10 ms
|
| 16 |
+
N_FFT = 512
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class LogMel(torch.nn.Module):
|
| 20 |
+
def __init__(self):
|
| 21 |
+
super().__init__()
|
| 22 |
+
self.mel = torchaudio.transforms.MelSpectrogram(
|
| 23 |
+
sample_rate=SAMPLE_RATE,
|
| 24 |
+
n_fft=N_FFT,
|
| 25 |
+
win_length=WIN_LENGTH,
|
| 26 |
+
hop_length=HOP_LENGTH,
|
| 27 |
+
n_mels=N_MELS,
|
| 28 |
+
center=True,
|
| 29 |
+
power=2.0,
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
EMA_ALPHA = 0.02 # ~0.5 s time constant at 10 ms hop
|
| 33 |
+
|
| 34 |
+
def forward(self, waveform):
|
| 35 |
+
"""waveform (B, samples) -> features (B, T, N_MELS).
|
| 36 |
+
|
| 37 |
+
Normalization is a causal per-channel EMA mean subtraction so the
|
| 38 |
+
exact same computation can run frame-by-frame on the device.
|
| 39 |
+
"""
|
| 40 |
+
mel = self.mel(waveform) # (B, n_mels, T)
|
| 41 |
+
logmel = torch.log(mel + 1e-6)
|
| 42 |
+
a = self.EMA_ALPHA
|
| 43 |
+
ema = torchaudio.functional.lfilter(
|
| 44 |
+
logmel,
|
| 45 |
+
a_coeffs=torch.tensor([1.0, -(1.0 - a)], device=logmel.device),
|
| 46 |
+
b_coeffs=torch.tensor([a, 0.0], device=logmel.device),
|
| 47 |
+
clamp=False,
|
| 48 |
+
)
|
| 49 |
+
logmel = logmel - ema
|
| 50 |
+
return logmel.transpose(1, 2) # (B, T, n_mels)
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def num_frames(num_samples):
|
| 54 |
+
"""Frame count produced for a waveform length (center=True)."""
|
| 55 |
+
return num_samples // HOP_LENGTH + 1
|
phoneme_engine/live_demo.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Live microphone wake word demo.
|
| 2 |
+
|
| 3 |
+
Type any phrase, it is enrolled instantly (no training) and spotted in the
|
| 4 |
+
live mic stream:
|
| 5 |
+
|
| 6 |
+
python -m phoneme_engine.live_demo "hey orbit"
|
| 7 |
+
|
| 8 |
+
Implementation notes:
|
| 9 |
+
- the causal model runs on a sliding 3 s window every 0.24 s; 0.24 s is
|
| 10 |
+
exactly 24 input frames = 12 output frames, so the emitted log-prob
|
| 11 |
+
stream stays frame-aligned (0.25 s = 12.5 frames caused stitching
|
| 12 |
+
artifacts and phantom detections)
|
| 13 |
+
- detections are energy-gated: the matched time span must contain real
|
| 14 |
+
audio well above the rolling noise floor, so silence/room noise cannot
|
| 15 |
+
fire no matter what the model hallucinates
|
| 16 |
+
- a mic level line is printed every few seconds for diagnosing input
|
| 17 |
+
device problems
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
import argparse
|
| 21 |
+
import queue
|
| 22 |
+
import sys
|
| 23 |
+
import time
|
| 24 |
+
from collections import deque
|
| 25 |
+
from pathlib import Path
|
| 26 |
+
|
| 27 |
+
import numpy as np
|
| 28 |
+
import sounddevice as sd
|
| 29 |
+
import torch
|
| 30 |
+
|
| 31 |
+
from .decoder import KeywordSpotter
|
| 32 |
+
from .features import HOP_LENGTH, SAMPLE_RATE, LogMel
|
| 33 |
+
from .model import PhonemeTCN
|
| 34 |
+
from .phones import ids_to_phones
|
| 35 |
+
|
| 36 |
+
ROOT = Path(__file__).resolve().parent.parent
|
| 37 |
+
WINDOW_S = 3.0
|
| 38 |
+
STEP_SAMPLES = 24 * HOP_LENGTH # 0.24 s = exactly 12 output frames
|
| 39 |
+
OUT_FRAMES_PER_STEP = 12
|
| 40 |
+
SAMPLES_PER_OUT_FRAME = 2 * HOP_LENGTH # 20 ms
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def dbfs(x):
|
| 44 |
+
rms = float(np.sqrt(np.mean(x ** 2)) + 1e-12)
|
| 45 |
+
return 20 * np.log10(rms)
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def main():
|
| 49 |
+
ap = argparse.ArgumentParser()
|
| 50 |
+
ap.add_argument("phrase")
|
| 51 |
+
ap.add_argument("--threshold", type=float, default=None,
|
| 52 |
+
help="avg per-frame deficit tolerance (more negative "
|
| 53 |
+
"= more lenient). Default: the enrollment's "
|
| 54 |
+
"auto-calibrated threshold, else -4.4")
|
| 55 |
+
ap.add_argument("--ckpt", default=str(ROOT / "checkpoints" / "best.pt"))
|
| 56 |
+
ap.add_argument("--show-phones", action="store_true",
|
| 57 |
+
help="print the greedy phoneme stream (debug)")
|
| 58 |
+
ap.add_argument("--min-speech-db", type=float, default=-50.0,
|
| 59 |
+
help="absolute floor: matched span must be louder")
|
| 60 |
+
ap.add_argument("--enrolled", action="store_true",
|
| 61 |
+
help="also match voice-enrolled pronunciation variants "
|
| 62 |
+
"(see phoneme_engine.enroll)")
|
| 63 |
+
ap.add_argument("--profile", default=None,
|
| 64 |
+
help="enrollment profile suffix, e.g. 'universal' "
|
| 65 |
+
"loads enrollments/<phrase>_universal.json")
|
| 66 |
+
args = ap.parse_args()
|
| 67 |
+
|
| 68 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 69 |
+
frontend = LogMel().to(device).eval()
|
| 70 |
+
model = PhonemeTCN().to(device).eval()
|
| 71 |
+
state = torch.load(args.ckpt, map_location=device, weights_only=True)
|
| 72 |
+
model.load_state_dict(state["model"])
|
| 73 |
+
per = state.get("best_per")
|
| 74 |
+
print(f"loaded {args.ckpt}" + (f" (dev PER {per:.3f})" if per else ""))
|
| 75 |
+
|
| 76 |
+
if args.enrolled:
|
| 77 |
+
from .enroll import load_spotter_configs
|
| 78 |
+
configs = load_spotter_configs(args.phrase, profile=args.profile)
|
| 79 |
+
else:
|
| 80 |
+
thr = args.threshold if args.threshold is not None else -4.4
|
| 81 |
+
configs = [(None, thr)]
|
| 82 |
+
spotters = []
|
| 83 |
+
for phones, thr in configs:
|
| 84 |
+
if args.threshold is not None:
|
| 85 |
+
thr = args.threshold
|
| 86 |
+
sp = KeywordSpotter(args.phrase, threshold=thr, phones=phones)
|
| 87 |
+
spotters.append(sp)
|
| 88 |
+
print(f"enrolled {args.phrase!r} -> {' '.join(sp.phones)} "
|
| 89 |
+
f"(threshold {thr:.2f})")
|
| 90 |
+
print("listening... Ctrl+C to stop")
|
| 91 |
+
|
| 92 |
+
q = queue.Queue()
|
| 93 |
+
|
| 94 |
+
def callback(indata, frames, t, status):
|
| 95 |
+
if status:
|
| 96 |
+
print(status, file=sys.stderr)
|
| 97 |
+
q.put(indata[:, 0].copy())
|
| 98 |
+
|
| 99 |
+
window = np.zeros(int(WINDOW_S * SAMPLE_RATE), dtype=np.float32)
|
| 100 |
+
pending = np.zeros(0, dtype=np.float32)
|
| 101 |
+
# per-output-frame energy history, aligned with spotter.t
|
| 102 |
+
energy_db = deque(maxlen=2000)
|
| 103 |
+
noise_floor = deque(maxlen=400) # ~8 s of frame energies
|
| 104 |
+
last_level_print = time.time()
|
| 105 |
+
|
| 106 |
+
with sd.InputStream(samplerate=SAMPLE_RATE, channels=1,
|
| 107 |
+
dtype="float32", blocksize=STEP_SAMPLES // 2,
|
| 108 |
+
callback=callback):
|
| 109 |
+
with torch.no_grad():
|
| 110 |
+
while True:
|
| 111 |
+
pending = np.concatenate([pending, q.get()])
|
| 112 |
+
while pending.shape[0] >= STEP_SAMPLES:
|
| 113 |
+
chunk = pending[:STEP_SAMPLES]
|
| 114 |
+
pending = pending[STEP_SAMPLES:]
|
| 115 |
+
window = np.concatenate([window[STEP_SAMPLES:], chunk])
|
| 116 |
+
|
| 117 |
+
for k in range(OUT_FRAMES_PER_STEP):
|
| 118 |
+
seg = chunk[k * SAMPLES_PER_OUT_FRAME:
|
| 119 |
+
(k + 1) * SAMPLES_PER_OUT_FRAME]
|
| 120 |
+
e = dbfs(seg)
|
| 121 |
+
energy_db.append(e)
|
| 122 |
+
noise_floor.append(e)
|
| 123 |
+
|
| 124 |
+
if time.time() - last_level_print > 5.0:
|
| 125 |
+
floor = np.percentile(noise_floor, 20)
|
| 126 |
+
print(f"[mic level {energy_db[-1]:6.1f} dBFS, "
|
| 127 |
+
f"noise floor {floor:6.1f}]", flush=True)
|
| 128 |
+
last_level_print = time.time()
|
| 129 |
+
|
| 130 |
+
wav = torch.from_numpy(window).unsqueeze(0).to(device)
|
| 131 |
+
logits = model(frontend(wav))
|
| 132 |
+
lp = torch.log_softmax(logits.float(), dim=2)[0]
|
| 133 |
+
lp_new = lp[-OUT_FRAMES_PER_STEP:].cpu().numpy()
|
| 134 |
+
if args.show_phones:
|
| 135 |
+
ids = lp_new.argmax(axis=1).tolist()
|
| 136 |
+
s = " ".join(ids_to_phones(ids))
|
| 137 |
+
if s.strip():
|
| 138 |
+
print(f"[{s}]", flush=True)
|
| 139 |
+
for t_idx in range(lp_new.shape[0]):
|
| 140 |
+
hit = None
|
| 141 |
+
for sp in spotters:
|
| 142 |
+
h = sp.step(lp_new[t_idx])
|
| 143 |
+
if h is not None and (hit is None
|
| 144 |
+
or h[0] > hit[0]):
|
| 145 |
+
hit = h
|
| 146 |
+
if hit is None:
|
| 147 |
+
continue
|
| 148 |
+
score, dur = hit
|
| 149 |
+
# one utterance shouldn't fire multiple variants
|
| 150 |
+
for sp in spotters:
|
| 151 |
+
sp.reset()
|
| 152 |
+
sp.cooldown = sp.refractory_frames
|
| 153 |
+
span = list(energy_db)[-(dur + (OUT_FRAMES_PER_STEP
|
| 154 |
+
- 1 - t_idx)):]
|
| 155 |
+
span = span[:dur] if dur <= len(span) else span
|
| 156 |
+
speech_db = float(np.mean(sorted(span)[len(span)//2:]))
|
| 157 |
+
floor = float(np.percentile(noise_floor, 20))
|
| 158 |
+
if speech_db < args.min_speech_db or \
|
| 159 |
+
speech_db < floor + 6.0:
|
| 160 |
+
print(f"(suppressed: score={score:.2f} but "
|
| 161 |
+
f"span energy {speech_db:.1f} dBFS ~ "
|
| 162 |
+
f"floor {floor:.1f})", flush=True)
|
| 163 |
+
continue
|
| 164 |
+
print(f"*** DETECTED {args.phrase!r} "
|
| 165 |
+
f"score={score:.2f} "
|
| 166 |
+
f"at {time.strftime('%H:%M:%S')} ***",
|
| 167 |
+
flush=True)
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
if __name__ == "__main__":
|
| 171 |
+
try:
|
| 172 |
+
main()
|
| 173 |
+
except KeyboardInterrupt:
|
| 174 |
+
print("\nstopped")
|
phoneme_engine/model.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Streaming phoneme recognizer: a small causal TCN with CTC output.
|
| 2 |
+
|
| 3 |
+
Design constraints (for later ESP32-S3 deployment via TFLite-Micro):
|
| 4 |
+
- only Conv1d / BatchNorm / ReLU / residual add (all INT8-friendly and
|
| 5 |
+
supported by esp-nn optimized kernels once folded/exported)
|
| 6 |
+
- strictly causal (left padding only) so it can run on a live stream
|
| 7 |
+
- ~350k parameters -> ~400 KB at INT8, fits the S3 easily
|
| 8 |
+
|
| 9 |
+
Input: (B, T, 40) log-mel frames at 10 ms
|
| 10 |
+
Output: (B, T', NUM_CLASSES) logits at 20 ms (stem has stride 2)
|
| 11 |
+
"""
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
import torch.nn as nn
|
| 15 |
+
|
| 16 |
+
from .phones import NUM_CLASSES
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class CausalConv1d(nn.Module):
|
| 20 |
+
"""Conv1d with left-only padding (streaming safe)."""
|
| 21 |
+
|
| 22 |
+
def __init__(self, in_ch, out_ch, kernel, stride=1, dilation=1, groups=1):
|
| 23 |
+
super().__init__()
|
| 24 |
+
self.left_pad = dilation * (kernel - 1)
|
| 25 |
+
self.conv = nn.Conv1d(in_ch, out_ch, kernel, stride=stride,
|
| 26 |
+
dilation=dilation, groups=groups)
|
| 27 |
+
|
| 28 |
+
def forward(self, x):
|
| 29 |
+
x = nn.functional.pad(x, (self.left_pad, 0))
|
| 30 |
+
return self.conv(x)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class TCNBlock(nn.Module):
|
| 34 |
+
"""Depthwise-separable causal conv block with residual."""
|
| 35 |
+
|
| 36 |
+
def __init__(self, channels, kernel, dilation):
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.dw = CausalConv1d(channels, channels, kernel,
|
| 39 |
+
dilation=dilation, groups=channels)
|
| 40 |
+
self.bn1 = nn.BatchNorm1d(channels)
|
| 41 |
+
self.pw = nn.Conv1d(channels, channels, 1)
|
| 42 |
+
self.bn2 = nn.BatchNorm1d(channels)
|
| 43 |
+
self.act = nn.ReLU()
|
| 44 |
+
|
| 45 |
+
def forward(self, x):
|
| 46 |
+
y = self.act(self.bn1(self.dw(x)))
|
| 47 |
+
y = self.bn2(self.pw(y))
|
| 48 |
+
return self.act(x + y)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class PhonemeTCN(nn.Module):
|
| 52 |
+
def __init__(self, n_mels=40, channels=192, kernel=5,
|
| 53 |
+
dilations=(1, 2, 4, 8, 1, 2, 4, 8)):
|
| 54 |
+
super().__init__()
|
| 55 |
+
self.stem = CausalConv1d(n_mels, channels, kernel, stride=2)
|
| 56 |
+
self.stem_bn = nn.BatchNorm1d(channels)
|
| 57 |
+
self.act = nn.ReLU()
|
| 58 |
+
self.blocks = nn.Sequential(
|
| 59 |
+
*[TCNBlock(channels, kernel, d) for d in dilations])
|
| 60 |
+
self.head = nn.Conv1d(channels, NUM_CLASSES, 1)
|
| 61 |
+
|
| 62 |
+
def forward(self, feats):
|
| 63 |
+
"""feats (B, T, n_mels) -> logits (B, T//2, NUM_CLASSES)."""
|
| 64 |
+
x = feats.transpose(1, 2) # (B, n_mels, T)
|
| 65 |
+
x = self.act(self.stem_bn(self.stem(x)))
|
| 66 |
+
x = self.blocks(x)
|
| 67 |
+
return self.head(x).transpose(1, 2) # (B, T', classes)
|
| 68 |
+
|
| 69 |
+
@staticmethod
|
| 70 |
+
def out_lengths(in_lengths):
|
| 71 |
+
"""Output frame count for input frame counts (stride-2 stem)."""
|
| 72 |
+
return torch.div(in_lengths - 1, 2, rounding_mode="floor") + 1
|
phoneme_engine/phones.py
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Phoneme inventory and lexicon for the phoneme engine.
|
| 2 |
+
|
| 3 |
+
Uses the 39-phone stress-less ARPAbet set (CMUdict) plus CTC blank at
|
| 4 |
+
index 0. Word pronunciations come from CMUdict, with grapheme-to-phoneme
|
| 5 |
+
fallback for out-of-vocabulary words (names, made-up wake words).
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
import functools
|
| 9 |
+
import os
|
| 10 |
+
import re
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
|
| 13 |
+
import pronouncing
|
| 14 |
+
|
| 15 |
+
# Pin NLTK data to a project-local folder. AppData can be virtualized
|
| 16 |
+
# differently per shell (app-container redirects), which made resources
|
| 17 |
+
# downloaded in one terminal invisible in another and tripped NLTK's
|
| 18 |
+
# path-security checks.
|
| 19 |
+
_NLTK_DIR = Path(__file__).resolve().parent.parent / "nltk_data"
|
| 20 |
+
os.environ["NLTK_DATA"] = str(_NLTK_DIR)
|
| 21 |
+
|
| 22 |
+
# 39 ARPAbet phones, stress stripped. Index 0 is reserved for CTC blank.
|
| 23 |
+
PHONES = [
|
| 24 |
+
"AA", "AE", "AH", "AO", "AW", "AY", "B", "CH", "D", "DH",
|
| 25 |
+
"EH", "ER", "EY", "F", "G", "HH", "IH", "IY", "JH", "K",
|
| 26 |
+
"L", "M", "N", "NG", "OW", "OY", "P", "R", "S", "SH",
|
| 27 |
+
"T", "TH", "UH", "UW", "V", "W", "Y", "Z", "ZH",
|
| 28 |
+
]
|
| 29 |
+
BLANK = 0
|
| 30 |
+
PHONE_TO_ID = {p: i + 1 for i, p in enumerate(PHONES)}
|
| 31 |
+
ID_TO_PHONE = {i + 1: p for i, p in enumerate(PHONES)}
|
| 32 |
+
NUM_CLASSES = len(PHONES) + 1 # 40 including blank
|
| 33 |
+
|
| 34 |
+
_g2p = None
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def _ensure_nltk_data():
|
| 38 |
+
import nltk
|
| 39 |
+
_NLTK_DIR.mkdir(exist_ok=True)
|
| 40 |
+
if str(_NLTK_DIR) not in nltk.data.path:
|
| 41 |
+
nltk.data.path.insert(0, str(_NLTK_DIR))
|
| 42 |
+
for res, sub in (("averaged_perceptron_tagger_eng", "taggers"),
|
| 43 |
+
("averaged_perceptron_tagger", "taggers"),
|
| 44 |
+
("cmudict", "corpora")):
|
| 45 |
+
try:
|
| 46 |
+
nltk.data.find(f"{sub}/{res}")
|
| 47 |
+
except LookupError:
|
| 48 |
+
nltk.download(res, download_dir=str(_NLTK_DIR), quiet=True)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def _g2p_phones(word):
|
| 52 |
+
global _g2p
|
| 53 |
+
if _g2p is None:
|
| 54 |
+
_ensure_nltk_data()
|
| 55 |
+
from g2p_en import G2p
|
| 56 |
+
_g2p = G2p()
|
| 57 |
+
return [re.sub(r"\d", "", p) for p in _g2p(word)
|
| 58 |
+
if re.match(r"^[A-Z]+$", re.sub(r"\d", "", p))]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
@functools.lru_cache(maxsize=200000)
|
| 62 |
+
def word_phones(word):
|
| 63 |
+
"""Phonemes for a single word (lowercase in, ARPAbet out)."""
|
| 64 |
+
prons = pronouncing.phones_for_word(word)
|
| 65 |
+
if prons:
|
| 66 |
+
return tuple(re.sub(r"\d", "", p) for p in prons[0].split())
|
| 67 |
+
return tuple(_g2p_phones(word))
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def text_to_phones(text):
|
| 71 |
+
"""Transcript -> flat phoneme list. Words it can't phonemize are skipped."""
|
| 72 |
+
phones = []
|
| 73 |
+
for word in re.findall(r"[a-z']+", text.lower()):
|
| 74 |
+
phones.extend(word_phones(word))
|
| 75 |
+
return phones
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def text_to_ids(text):
|
| 79 |
+
return [PHONE_TO_ID[p] for p in text_to_phones(text) if p in PHONE_TO_ID]
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def ids_to_phones(ids):
|
| 83 |
+
return [ID_TO_PHONE[i] for i in ids if i in ID_TO_PHONE]
|
phoneme_engine/quantize.py
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Post-training INT8 quantization for the PhonemeTCN, targeting a small
|
| 2 |
+
custom streaming engine on the ESP32-S3.
|
| 3 |
+
|
| 4 |
+
Scheme (classic TFLite-style, but explicit and portable):
|
| 5 |
+
- weights: symmetric per-output-channel int8, BN folded into conv
|
| 6 |
+
- activations: per-tensor asymmetric-free symmetric int8 (ReLU outputs
|
| 7 |
+
use unsigned range via zero offset 0..127 semantics kept simple:
|
| 8 |
+
symmetric [-127,127] everywhere)
|
| 9 |
+
- accumulators: int32; requantization by fixed-point multiplier per layer
|
| 10 |
+
|
| 11 |
+
Steps:
|
| 12 |
+
python -m phoneme_engine.quantize calibrate # activation ranges
|
| 13 |
+
python -m phoneme_engine.quantize verify # int8 sim vs float parity
|
| 14 |
+
python -m phoneme_engine.quantize export # C header + weights bin
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
import json
|
| 18 |
+
import sys
|
| 19 |
+
from pathlib import Path
|
| 20 |
+
|
| 21 |
+
import numpy as np
|
| 22 |
+
import torch
|
| 23 |
+
|
| 24 |
+
from .data import LibriPhonemes, collate
|
| 25 |
+
from .features import LogMel
|
| 26 |
+
from .model import PhonemeTCN
|
| 27 |
+
|
| 28 |
+
ROOT = Path(__file__).resolve().parent.parent
|
| 29 |
+
QDIR = ROOT / "export"
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def fold_bn(conv_w, conv_b, bn):
|
| 33 |
+
"""Fold BatchNorm into conv weights/bias (eval-mode running stats)."""
|
| 34 |
+
gamma = bn.weight.detach().numpy()
|
| 35 |
+
beta = bn.bias.detach().numpy()
|
| 36 |
+
mean = bn.running_mean.detach().numpy()
|
| 37 |
+
var = bn.running_var.detach().numpy()
|
| 38 |
+
scale = gamma / np.sqrt(var + bn.eps)
|
| 39 |
+
w = conv_w * scale[:, None, None]
|
| 40 |
+
b = (conv_b if conv_b is not None else 0.0) * scale + beta - mean * scale
|
| 41 |
+
return w, b
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def extract_layers(model):
|
| 45 |
+
"""Flatten the model into a list of layer dicts with folded BN.
|
| 46 |
+
|
| 47 |
+
Layer kinds: conv (dense conv1d), dw (depthwise), each with
|
| 48 |
+
weight (out, in, k) / (ch, 1, k), bias, stride, dilation, plus
|
| 49 |
+
residual bookkeeping: blocks add their input to the pw output.
|
| 50 |
+
"""
|
| 51 |
+
layers = []
|
| 52 |
+
w, b = fold_bn(model.stem.conv.weight.detach().numpy(),
|
| 53 |
+
None if model.stem.conv.bias is None
|
| 54 |
+
else model.stem.conv.bias.detach().numpy(),
|
| 55 |
+
model.stem_bn)
|
| 56 |
+
layers.append(dict(kind="conv", w=w, b=b, stride=2, dilation=1,
|
| 57 |
+
relu=True, residual=False))
|
| 58 |
+
for blk in model.blocks:
|
| 59 |
+
w, b = fold_bn(blk.dw.conv.weight.detach().numpy(),
|
| 60 |
+
None if blk.dw.conv.bias is None
|
| 61 |
+
else blk.dw.conv.bias.detach().numpy(),
|
| 62 |
+
blk.bn1)
|
| 63 |
+
layers.append(dict(kind="dw", w=w, b=b, stride=1,
|
| 64 |
+
dilation=blk.dw.conv.dilation[0], relu=True,
|
| 65 |
+
residual=False))
|
| 66 |
+
w, b = fold_bn(blk.pw.weight.detach().numpy(),
|
| 67 |
+
None if blk.pw.bias is None
|
| 68 |
+
else blk.pw.bias.detach().numpy(),
|
| 69 |
+
blk.bn2)
|
| 70 |
+
# pw output adds the block input, then ReLU
|
| 71 |
+
layers.append(dict(kind="conv", w=w, b=b, stride=1, dilation=1,
|
| 72 |
+
relu=True, residual=True))
|
| 73 |
+
layers.append(dict(kind="conv",
|
| 74 |
+
w=model.head.weight.detach().numpy(),
|
| 75 |
+
b=model.head.bias.detach().numpy(),
|
| 76 |
+
stride=1, dilation=1, relu=False, residual=False))
|
| 77 |
+
return layers
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def float_forward(layers, feats):
|
| 81 |
+
"""Reference float forward pass on (T, 40) features using the flat
|
| 82 |
+
layer list. Must match the PyTorch model exactly (verified)."""
|
| 83 |
+
x = feats.T # (C, T)
|
| 84 |
+
block_input = None
|
| 85 |
+
for lay in layers:
|
| 86 |
+
if lay["kind"] == "dw":
|
| 87 |
+
block_input = x # residual adds the block's input, saved here
|
| 88 |
+
w, b = lay["w"], lay["b"]
|
| 89 |
+
k = w.shape[2]
|
| 90 |
+
d = lay["dilation"]
|
| 91 |
+
pad = d * (k - 1)
|
| 92 |
+
xin = np.pad(x, ((0, 0), (pad, 0)))
|
| 93 |
+
T = x.shape[1]
|
| 94 |
+
out_T = (T - 1) // lay["stride"] + 1
|
| 95 |
+
out_C = w.shape[0]
|
| 96 |
+
y = np.zeros((out_C, out_T), dtype=np.float64)
|
| 97 |
+
for t in range(out_T):
|
| 98 |
+
base = t * lay["stride"] + pad
|
| 99 |
+
taps = xin[:, [base - d * (k - 1 - i) for i in range(k)]]
|
| 100 |
+
if lay["kind"] == "dw":
|
| 101 |
+
y[:, t] = (taps * w[:, 0, :]).sum(axis=1) + b
|
| 102 |
+
else:
|
| 103 |
+
y[:, t] = np.tensordot(w, taps, axes=([1, 2], [0, 1])) + b
|
| 104 |
+
if lay["residual"]:
|
| 105 |
+
y = y + block_input[:, :out_T]
|
| 106 |
+
if lay["relu"]:
|
| 107 |
+
y = np.maximum(y, 0.0)
|
| 108 |
+
x = y
|
| 109 |
+
return x.T # (T, classes)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def collect_calibration_feats(n_utts=64):
|
| 113 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 114 |
+
frontend = LogMel().to(device).eval()
|
| 115 |
+
ds = LibriPhonemes(str(ROOT / "data"), "dev-clean")
|
| 116 |
+
feats = []
|
| 117 |
+
with torch.no_grad():
|
| 118 |
+
for i in range(0, n_utts * 40, 40):
|
| 119 |
+
wav, _ = ds[i % len(ds)]
|
| 120 |
+
f = frontend(wav.unsqueeze(0).to(device))[0].cpu().numpy()
|
| 121 |
+
feats.append(f[:400]) # up to 4 s per utterance
|
| 122 |
+
if len(feats) >= n_utts:
|
| 123 |
+
break
|
| 124 |
+
return feats
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def calibrate_scales(layers, feats_list, pctl=99.95):
|
| 128 |
+
"""Per-layer activation scales from representative audio: runs the
|
| 129 |
+
float replay and records robust max-abs at every layer boundary."""
|
| 130 |
+
n_layers = len(layers)
|
| 131 |
+
maxima = [[] for _ in range(n_layers + 1)] # +1 for the input feats
|
| 132 |
+
for feats in feats_list:
|
| 133 |
+
maxima[0].append(np.percentile(np.abs(feats), pctl))
|
| 134 |
+
x = feats.T
|
| 135 |
+
block_input = None
|
| 136 |
+
for li, lay in enumerate(layers):
|
| 137 |
+
if lay["kind"] == "dw":
|
| 138 |
+
block_input = x
|
| 139 |
+
w, b = lay["w"], lay["b"]
|
| 140 |
+
k = w.shape[2]
|
| 141 |
+
d = lay["dilation"]
|
| 142 |
+
pad = d * (k - 1)
|
| 143 |
+
xin = np.pad(x, ((0, 0), (pad, 0)))
|
| 144 |
+
out_T = (x.shape[1] - 1) // lay["stride"] + 1
|
| 145 |
+
y = np.zeros((w.shape[0], out_T))
|
| 146 |
+
for t in range(out_T):
|
| 147 |
+
base = t * lay["stride"] + pad
|
| 148 |
+
taps = xin[:, [base - d * (k - 1 - i) for i in range(k)]]
|
| 149 |
+
if lay["kind"] == "dw":
|
| 150 |
+
y[:, t] = (taps * w[:, 0, :]).sum(axis=1) + b
|
| 151 |
+
else:
|
| 152 |
+
y[:, t] = np.tensordot(w, taps,
|
| 153 |
+
axes=([1, 2], [0, 1])) + b
|
| 154 |
+
if lay["residual"]:
|
| 155 |
+
y = y + block_input[:, :out_T]
|
| 156 |
+
if lay["relu"]:
|
| 157 |
+
y = np.maximum(y, 0.0)
|
| 158 |
+
maxima[li + 1].append(np.percentile(np.abs(y), pctl))
|
| 159 |
+
x = y
|
| 160 |
+
scales = [max(float(np.max(m)), 1e-3) / 127.0 for m in maxima]
|
| 161 |
+
return scales
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
def quantize_weights(layers):
|
| 165 |
+
"""Symmetric per-output-channel int8 weights; returns quantized copies
|
| 166 |
+
(float values on the int8 grid) plus the raw int8 arrays and scales."""
|
| 167 |
+
qlayers = []
|
| 168 |
+
for lay in layers:
|
| 169 |
+
w = lay["w"]
|
| 170 |
+
s_w = np.abs(w).reshape(w.shape[0], -1).max(axis=1) / 127.0
|
| 171 |
+
s_w = np.maximum(s_w, 1e-8)
|
| 172 |
+
w_int = np.clip(np.round(w / s_w[:, None, None]), -127, 127)
|
| 173 |
+
q = dict(lay)
|
| 174 |
+
q["w"] = w_int * s_w[:, None, None]
|
| 175 |
+
q["w_int"] = w_int.astype(np.int8)
|
| 176 |
+
q["s_w"] = s_w
|
| 177 |
+
qlayers.append(q)
|
| 178 |
+
return qlayers
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
def fake_quant_forward(qlayers, scales, feats):
|
| 182 |
+
"""Float replay with activations snapped to the int8 grid at every
|
| 183 |
+
layer boundary -- numerically equivalent to the integer engine."""
|
| 184 |
+
def snap(x, s):
|
| 185 |
+
return np.clip(np.round(x / s), -127, 127) * s
|
| 186 |
+
|
| 187 |
+
x = snap(feats.T, scales[0])
|
| 188 |
+
block_input = None
|
| 189 |
+
block_input_scale = None
|
| 190 |
+
for li, lay in enumerate(qlayers):
|
| 191 |
+
if lay["kind"] == "dw":
|
| 192 |
+
block_input = x
|
| 193 |
+
block_input_scale = scales[li]
|
| 194 |
+
w, b = lay["w"], lay["b"]
|
| 195 |
+
k = w.shape[2]
|
| 196 |
+
d = lay["dilation"]
|
| 197 |
+
pad = d * (k - 1)
|
| 198 |
+
xin = np.pad(x, ((0, 0), (pad, 0)))
|
| 199 |
+
out_T = (x.shape[1] - 1) // lay["stride"] + 1
|
| 200 |
+
y = np.zeros((w.shape[0], out_T))
|
| 201 |
+
for t in range(out_T):
|
| 202 |
+
base = t * lay["stride"] + pad
|
| 203 |
+
taps = xin[:, [base - d * (k - 1 - i) for i in range(k)]]
|
| 204 |
+
if lay["kind"] == "dw":
|
| 205 |
+
y[:, t] = (taps * w[:, 0, :]).sum(axis=1) + b
|
| 206 |
+
else:
|
| 207 |
+
y[:, t] = np.tensordot(w, taps, axes=([1, 2], [0, 1])) + b
|
| 208 |
+
if lay["residual"]:
|
| 209 |
+
y = y + snap(block_input[:, :out_T], block_input_scale)
|
| 210 |
+
if lay["relu"]:
|
| 211 |
+
y = np.maximum(y, 0.0)
|
| 212 |
+
x = snap(y, scales[li + 1])
|
| 213 |
+
return x.T
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def main():
|
| 217 |
+
cmd = sys.argv[1] if len(sys.argv) > 1 else "verify"
|
| 218 |
+
QDIR.mkdir(exist_ok=True)
|
| 219 |
+
device = "cpu"
|
| 220 |
+
model = PhonemeTCN().eval()
|
| 221 |
+
state = torch.load(ROOT / "checkpoints" / "best.pt",
|
| 222 |
+
map_location=device, weights_only=True)
|
| 223 |
+
model.load_state_dict(state["model"])
|
| 224 |
+
layers = extract_layers(model)
|
| 225 |
+
print(f"{len(layers)} layers extracted")
|
| 226 |
+
|
| 227 |
+
if cmd == "sanity":
|
| 228 |
+
feats = collect_calibration_feats(2)
|
| 229 |
+
x = torch.from_numpy(feats[0]).float().unsqueeze(0)
|
| 230 |
+
with torch.no_grad():
|
| 231 |
+
ref = model(x)[0].numpy()
|
| 232 |
+
ours = float_forward(layers, feats[0])
|
| 233 |
+
err = np.abs(ref - ours).max()
|
| 234 |
+
print(f"float reference vs pytorch max abs err: {err:.2e}")
|
| 235 |
+
assert err < 1e-3, "layer extraction is wrong"
|
| 236 |
+
print("sanity OK")
|
| 237 |
+
|
| 238 |
+
elif cmd == "calibrate":
|
| 239 |
+
feats_list = collect_calibration_feats(48)
|
| 240 |
+
print(f"calibrating on {len(feats_list)} utterances...")
|
| 241 |
+
scales = calibrate_scales(layers, feats_list)
|
| 242 |
+
(QDIR / "act_scales.json").write_text(json.dumps(scales))
|
| 243 |
+
print("activation scales:", [round(s, 5) for s in scales])
|
| 244 |
+
print(f"saved -> {QDIR / 'act_scales.json'}")
|
| 245 |
+
|
| 246 |
+
elif cmd == "verify":
|
| 247 |
+
scales = json.loads((QDIR / "act_scales.json").read_text())
|
| 248 |
+
qlayers = quantize_weights(layers)
|
| 249 |
+
from .decoder import KeywordSpotter
|
| 250 |
+
from .spot_file import load_wav
|
| 251 |
+
from .features import LogMel
|
| 252 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 253 |
+
frontend = LogMel().to(device).eval()
|
| 254 |
+
|
| 255 |
+
agree = tot = 0
|
| 256 |
+
deltas = []
|
| 257 |
+
clips = [(f, "sakura") for f in
|
| 258 |
+
sorted((ROOT / "test_clips" / "sakura_piper").glob("*.wav"))[:8]]
|
| 259 |
+
clips += [(f, "hey orbit") for f in
|
| 260 |
+
sorted((ROOT / "test_clips" / "hey_orbit_piper").glob("*.wav"))[:8]]
|
| 261 |
+
clips.append((ROOT / "diag_last.wav", "hey orbit"))
|
| 262 |
+
for f, phrase in clips:
|
| 263 |
+
wav = load_wav(str(f), device)
|
| 264 |
+
with torch.no_grad():
|
| 265 |
+
feats = frontend(wav)[0].cpu().numpy()
|
| 266 |
+
lf = float_forward(layers, feats)
|
| 267 |
+
lq = fake_quant_forward(qlayers, scales, feats)
|
| 268 |
+
agree += (lf.argmax(1) == lq.argmax(1)).sum()
|
| 269 |
+
tot += lf.shape[0]
|
| 270 |
+
# spotting parity: decoder consumes logit differences directly
|
| 271 |
+
sp = KeywordSpotter(phrase)
|
| 272 |
+
sf_ = sp.best_score(lf - lf.max(axis=1, keepdims=True))
|
| 273 |
+
sq_ = sp.best_score(lq - lq.max(axis=1, keepdims=True))
|
| 274 |
+
if np.isfinite(sf_) or np.isfinite(sq_):
|
| 275 |
+
deltas.append(abs(sf_ - sq_))
|
| 276 |
+
print(f" {f.name} [{phrase}]: float {sf_:7.2f} int8 {sq_:7.2f}")
|
| 277 |
+
print(f"frame argmax agreement: {agree/tot:.4f}")
|
| 278 |
+
finite = [d for d in deltas if np.isfinite(d)]
|
| 279 |
+
print(f"spot score delta (finite pairs): max {max(finite):.3f} "
|
| 280 |
+
f"mean {np.mean(finite):.3f}; gate flips: "
|
| 281 |
+
f"{len(deltas) - len(finite)}")
|
| 282 |
+
|
| 283 |
+
elif cmd == "export":
|
| 284 |
+
scales = json.loads((QDIR / "act_scales.json").read_text())
|
| 285 |
+
qlayers = quantize_weights(layers)
|
| 286 |
+
from .phones import PHONES
|
| 287 |
+
lines = ["// Auto-generated by phoneme_engine.quantize export",
|
| 288 |
+
"// PhonemeTCN int8 weights + scales for the streaming",
|
| 289 |
+
"// wake word engine. Do not edit by hand.",
|
| 290 |
+
"#pragma once", "#include <stdint.h>", ""]
|
| 291 |
+
lines.append(f"#define PWW_NUM_LAYERS {len(qlayers)}")
|
| 292 |
+
lines.append(f"#define PWW_NUM_CLASSES {len(PHONES) + 1}")
|
| 293 |
+
lines.append(f"#define PWW_INPUT_SCALE {scales[0]:.8f}f")
|
| 294 |
+
lines.append(f"#define PWW_LOGIT_SCALE {scales[-1]:.8f}f")
|
| 295 |
+
lines.append("")
|
| 296 |
+
phones_str = ", ".join(f'"{p}"' for p in PHONES)
|
| 297 |
+
lines.append(f"static const char *PWW_PHONES[] = {{{phones_str}}};")
|
| 298 |
+
lines.append("")
|
| 299 |
+
meta_rows = []
|
| 300 |
+
total_bytes = 0
|
| 301 |
+
for li, lay in enumerate(qlayers):
|
| 302 |
+
w_int = lay["w_int"]
|
| 303 |
+
out_c, in_c, k = w_int.shape
|
| 304 |
+
flat = w_int.flatten()
|
| 305 |
+
total_bytes += flat.size
|
| 306 |
+
arr = ", ".join(str(int(v)) for v in flat)
|
| 307 |
+
lines.append(f"static const int8_t PWW_W{li}[] = {{{arr}}};")
|
| 308 |
+
# combined scale per channel: s_in * s_w[c] (float requant)
|
| 309 |
+
comb = scales[li] * lay["s_w"]
|
| 310 |
+
arr = ", ".join(f"{v:.8e}f" for v in comb)
|
| 311 |
+
lines.append(f"static const float PWW_S{li}[] = {{{arr}}};")
|
| 312 |
+
arr = ", ".join(f"{v:.8e}f" for v in lay["b"])
|
| 313 |
+
lines.append(f"static const float PWW_B{li}[] = {{{arr}}};")
|
| 314 |
+
# depthwise weights are stored (ch, 1, k): the layer's true
|
| 315 |
+
# input width is out_c, not the stored dim
|
| 316 |
+
eff_in = out_c if lay["kind"] == "dw" else in_c
|
| 317 |
+
meta_rows.append(
|
| 318 |
+
f" {{{1 if lay['kind'] == 'dw' else 0}, {eff_in}, {out_c}, "
|
| 319 |
+
f"{k}, {lay['stride']}, {lay['dilation']}, "
|
| 320 |
+
f"{1 if lay['relu'] else 0}, {1 if lay['residual'] else 0}, "
|
| 321 |
+
f"PWW_W{li}, PWW_S{li}, PWW_B{li}, "
|
| 322 |
+
f"{scales[li + 1]:.8f}f}}")
|
| 323 |
+
lines.append("")
|
| 324 |
+
lines.append(
|
| 325 |
+
"typedef struct { uint8_t is_dw; uint16_t in_c, out_c; "
|
| 326 |
+
"uint8_t k, stride, dilation, relu, residual; "
|
| 327 |
+
"const int8_t *w; const float *s; const float *b; "
|
| 328 |
+
"float out_scale; } pww_layer_t;")
|
| 329 |
+
lines.append("")
|
| 330 |
+
lines.append("static const pww_layer_t PWW_LAYERS[] = {")
|
| 331 |
+
lines.append(",\n".join(meta_rows))
|
| 332 |
+
lines.append("};")
|
| 333 |
+
path = QDIR / "model_int8.h"
|
| 334 |
+
path.write_text("\n".join(lines), encoding="utf-8")
|
| 335 |
+
print(f"exported {total_bytes/1024:.0f} KB of int8 weights "
|
| 336 |
+
f"-> {path}")
|
| 337 |
+
|
| 338 |
+
elif cmd == "export-frontend":
|
| 339 |
+
# exact DSP constants from the training frontend, so the C mel
|
| 340 |
+
# frontend is identical by construction
|
| 341 |
+
from .features import (HOP_LENGTH, N_FFT, N_MELS, SAMPLE_RATE,
|
| 342 |
+
WIN_LENGTH, LogMel)
|
| 343 |
+
fe = LogMel()
|
| 344 |
+
fb = fe.mel.mel_scale.fb.numpy() # (n_freqs, n_mels)
|
| 345 |
+
win = torch.hann_window(WIN_LENGTH, periodic=True).numpy()
|
| 346 |
+
lines = ["// Auto-generated: mel frontend constants (exact copy of",
|
| 347 |
+
"// the training features). Do not edit.",
|
| 348 |
+
"#pragma once", ""]
|
| 349 |
+
lines.append(f"#define PWW_FE_SR {SAMPLE_RATE}")
|
| 350 |
+
lines.append(f"#define PWW_FE_NFFT {N_FFT}")
|
| 351 |
+
lines.append(f"#define PWW_FE_WIN {WIN_LENGTH}")
|
| 352 |
+
lines.append(f"#define PWW_FE_HOP {HOP_LENGTH}")
|
| 353 |
+
lines.append(f"#define PWW_FE_NMELS {N_MELS}")
|
| 354 |
+
lines.append(f"#define PWW_FE_NFREQS {fb.shape[0]}")
|
| 355 |
+
lines.append("#define PWW_FE_EMA_ALPHA 0.02f")
|
| 356 |
+
lines.append("")
|
| 357 |
+
arr = ", ".join(f"{v:.8e}f" for v in win)
|
| 358 |
+
lines.append(f"static const float PWW_FE_HANN[] = {{{arr}}};")
|
| 359 |
+
arr = ", ".join(f"{v:.8e}f" for v in fb.T.flatten())
|
| 360 |
+
lines.append("// mel filterbank, row-major (n_mels, n_freqs)")
|
| 361 |
+
lines.append(f"static const float PWW_FE_MELFB[] = {{{arr}}};")
|
| 362 |
+
path = QDIR / "frontend_data.h"
|
| 363 |
+
path.write_text("\n".join(lines), encoding="utf-8")
|
| 364 |
+
print(f"exported frontend constants -> {path}")
|
| 365 |
+
|
| 366 |
+
|
| 367 |
+
if __name__ == "__main__":
|
| 368 |
+
main()
|
phoneme_engine/score_phrase.py
ADDED
|
@@ -0,0 +1,340 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Wake word phrase scorer.
|
| 2 |
+
|
| 3 |
+
Scores a candidate wake word phrase on how resistant it will be to false
|
| 4 |
+
activations, using phoneme statistics of spoken English:
|
| 5 |
+
|
| 6 |
+
1. rarity - mean surprisal (-log2 P) of the phrase's phoneme bigrams,
|
| 7 |
+
estimated from CMUdict pronunciations weighted by real
|
| 8 |
+
word-usage frequency (wordfreq). Rare sound sequences are
|
| 9 |
+
less likely to appear in everyday speech.
|
| 10 |
+
2. confusability - minimum weighted phonetic edit distance (normalized)
|
| 11 |
+
between the phrase and common words / common two-word
|
| 12 |
+
concatenations. Low distance = easy to false-trigger.
|
| 13 |
+
3. syllables - counts vowel nuclei; 3-5 syllables is the sweet spot for
|
| 14 |
+
small streaming models like microWakeWord.
|
| 15 |
+
4. diversity - fraction of distinct phonemes; repeated sounds carry less
|
| 16 |
+
discriminative information.
|
| 17 |
+
|
| 18 |
+
Usage:
|
| 19 |
+
python score_phrase.py "hey jarvis" # score one phrase
|
| 20 |
+
python score_phrase.py --rank candidates.txt # rank a file of phrases
|
| 21 |
+
python score_phrase.py --suggest 8 # auto-generate + rank candidates
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
import argparse
|
| 25 |
+
import functools
|
| 26 |
+
import json
|
| 27 |
+
import math
|
| 28 |
+
import random
|
| 29 |
+
import re
|
| 30 |
+
import sys
|
| 31 |
+
|
| 32 |
+
import pronouncing
|
| 33 |
+
from wordfreq import top_n_list, word_frequency
|
| 34 |
+
|
| 35 |
+
VOWELS = {
|
| 36 |
+
"AA", "AE", "AH", "AO", "AW", "AY", "EH", "ER", "EY",
|
| 37 |
+
"IH", "IY", "OW", "OY", "UH", "UW",
|
| 38 |
+
}
|
| 39 |
+
|
| 40 |
+
# Broad phonetic classes used to soften substitution costs: confusing two
|
| 41 |
+
# sounds from the same class is cheaper (more likely for a model) than
|
| 42 |
+
# crossing classes.
|
| 43 |
+
PHONE_CLASS = {}
|
| 44 |
+
for _class, _phones in {
|
| 45 |
+
"vowel": VOWELS,
|
| 46 |
+
"stop": {"B", "D", "G", "K", "P", "T"},
|
| 47 |
+
"affricate": {"CH", "JH"},
|
| 48 |
+
"fricative": {"DH", "F", "HH", "S", "SH", "TH", "V", "Z", "ZH"},
|
| 49 |
+
"nasal": {"M", "N", "NG"},
|
| 50 |
+
"liquid": {"L", "R"},
|
| 51 |
+
"glide": {"W", "Y"},
|
| 52 |
+
}.items():
|
| 53 |
+
for _p in _phones:
|
| 54 |
+
PHONE_CLASS[_p] = _class
|
| 55 |
+
|
| 56 |
+
# Acoustic reliability of each phone for a small on-device model,
|
| 57 |
+
# validated empirically with live tests: V/F/TH-type weak fricatives get
|
| 58 |
+
# swallowed or confused; nasals, stops and sibilants are heard stably.
|
| 59 |
+
RELIABILITY = {}
|
| 60 |
+
for _score, _phones in [
|
| 61 |
+
(1.0, ["N", "M", "K", "T", "D", "P", "B", "G", "S", "SH", "CH", "JH",
|
| 62 |
+
"NG", "Z"]),
|
| 63 |
+
(0.8, ["AA", "AE", "AH", "AO", "AW", "AY", "EH", "ER", "EY", "IH",
|
| 64 |
+
"IY", "OW", "OY", "UH", "UW", "L", "R"]),
|
| 65 |
+
(0.5, ["F", "V", "W", "Y", "HH", "TH", "DH", "ZH"]),
|
| 66 |
+
]:
|
| 67 |
+
for _p in _phones:
|
| 68 |
+
RELIABILITY[_p] = _score
|
| 69 |
+
|
| 70 |
+
_g2p = None
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def g2p_word(word):
|
| 74 |
+
"""Phonemes for an out-of-vocabulary word via letter-to-sound rules."""
|
| 75 |
+
global _g2p
|
| 76 |
+
if _g2p is None:
|
| 77 |
+
from g2p_en import G2p
|
| 78 |
+
_g2p = G2p()
|
| 79 |
+
return [re.sub(r"\d", "", p) for p in _g2p(word) if re.match(r"[A-Z]", p)]
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def phones_for_word(word):
|
| 83 |
+
"""CMUdict pronunciation if available, else grapheme-to-phoneme."""
|
| 84 |
+
prons = pronouncing.phones_for_word(word.lower())
|
| 85 |
+
if prons:
|
| 86 |
+
return re.sub(r"\d", "", prons[0]).split()
|
| 87 |
+
return g2p_word(word)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def phones_for_phrase(phrase):
|
| 91 |
+
phones = []
|
| 92 |
+
for word in re.findall(r"[a-zA-Z']+", phrase):
|
| 93 |
+
phones.extend(phones_for_word(word))
|
| 94 |
+
return phones
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
@functools.lru_cache(maxsize=1)
|
| 98 |
+
def bigram_model(vocab_size=30000):
|
| 99 |
+
"""Frequency-weighted phoneme bigram probabilities for spoken English.
|
| 100 |
+
|
| 101 |
+
'#' marks a word boundary so cross-word statistics are approximated.
|
| 102 |
+
Returns (probabilities, common_word_phones) where common_word_phones is
|
| 103 |
+
the pronunciation cache used by the confusability search.
|
| 104 |
+
"""
|
| 105 |
+
counts = {}
|
| 106 |
+
total = 0.0
|
| 107 |
+
word_phones = {}
|
| 108 |
+
for word in top_n_list("en", vocab_size):
|
| 109 |
+
prons = pronouncing.phones_for_word(word)
|
| 110 |
+
if not prons:
|
| 111 |
+
continue
|
| 112 |
+
phones = re.sub(r"\d", "", prons[0]).split()
|
| 113 |
+
if not phones:
|
| 114 |
+
continue
|
| 115 |
+
word_phones[word] = phones
|
| 116 |
+
weight = word_frequency(word, "en")
|
| 117 |
+
seq = ["#"] + phones + ["#"]
|
| 118 |
+
for a, b in zip(seq, seq[1:]):
|
| 119 |
+
counts[(a, b)] = counts.get((a, b), 0.0) + weight
|
| 120 |
+
total += weight
|
| 121 |
+
n_types = max(len(counts), 1)
|
| 122 |
+
probs = {k: v / total for k, v in counts.items()}
|
| 123 |
+
floor = 0.01 / (total * n_types) # smoothing for unseen bigrams
|
| 124 |
+
return probs, floor, word_phones
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def rarity(phones):
|
| 128 |
+
"""Mean surprisal in bits of the phrase's phoneme bigrams."""
|
| 129 |
+
probs, floor, _ = bigram_model()
|
| 130 |
+
seq = ["#"] + phones + ["#"]
|
| 131 |
+
bits = [-math.log2(probs.get((a, b), floor)) for a, b in zip(seq, seq[1:])]
|
| 132 |
+
return sum(bits) / len(bits)
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def sub_cost(a, b):
|
| 136 |
+
if a == b:
|
| 137 |
+
return 0.0
|
| 138 |
+
if PHONE_CLASS.get(a) == PHONE_CLASS.get(b):
|
| 139 |
+
return 0.5
|
| 140 |
+
return 1.0
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
def edit_distance(a, b, bail_above=None):
|
| 144 |
+
"""Weighted phonetic Levenshtein distance with optional early exit."""
|
| 145 |
+
prev = [i * 1.0 for i in range(len(b) + 1)]
|
| 146 |
+
for i, pa in enumerate(a, 1):
|
| 147 |
+
cur = [i * 1.0]
|
| 148 |
+
for j, pb in enumerate(b, 1):
|
| 149 |
+
cur.append(min(prev[j] + 1.0, cur[j - 1] + 1.0,
|
| 150 |
+
prev[j - 1] + sub_cost(pa, pb)))
|
| 151 |
+
if bail_above is not None and min(cur) > bail_above:
|
| 152 |
+
return bail_above + 1.0
|
| 153 |
+
prev = cur
|
| 154 |
+
return prev[-1]
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def confusability(phones, n_single=3000, n_pairs=400):
|
| 158 |
+
"""Min normalized phonetic distance to common words and word pairs.
|
| 159 |
+
|
| 160 |
+
Higher = safer. Distances are normalized by the longer sequence length so
|
| 161 |
+
short and long phrases are comparable.
|
| 162 |
+
"""
|
| 163 |
+
_, _, word_phones = bigram_model()
|
| 164 |
+
singles = [w for w in top_n_list("en", n_single) if w in word_phones]
|
| 165 |
+
pair_words = [w for w in singles[:n_pairs]]
|
| 166 |
+
|
| 167 |
+
best = float("inf")
|
| 168 |
+
best_match = None
|
| 169 |
+
for w in singles:
|
| 170 |
+
wp = word_phones[w]
|
| 171 |
+
norm = max(len(phones), len(wp))
|
| 172 |
+
d = edit_distance(phones, wp, bail_above=best * norm) / norm
|
| 173 |
+
if d < best:
|
| 174 |
+
best, best_match = d, w
|
| 175 |
+
|
| 176 |
+
# Two-word concatenations approximate the phrase appearing inside
|
| 177 |
+
# continuous speech. Prune by length: only pairs within +/-4 phones.
|
| 178 |
+
target_len = len(phones)
|
| 179 |
+
for w1 in pair_words:
|
| 180 |
+
p1 = word_phones[w1]
|
| 181 |
+
if len(p1) >= target_len + 4:
|
| 182 |
+
continue
|
| 183 |
+
for w2 in pair_words:
|
| 184 |
+
cat = p1 + word_phones[w2]
|
| 185 |
+
if abs(len(cat) - target_len) > 4:
|
| 186 |
+
continue
|
| 187 |
+
norm = max(target_len, len(cat))
|
| 188 |
+
d = edit_distance(phones, cat, bail_above=best * norm) / norm
|
| 189 |
+
if d < best:
|
| 190 |
+
best, best_match = d, f"{w1} {w2}"
|
| 191 |
+
return best, best_match
|
| 192 |
+
|
| 193 |
+
|
| 194 |
+
def syllable_count(phones):
|
| 195 |
+
return sum(1 for p in phones if p in VOWELS)
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
def score_phrase(phrase, fast=False):
|
| 199 |
+
phones = phones_for_phrase(phrase)
|
| 200 |
+
if not phones:
|
| 201 |
+
raise ValueError(f"could not phonemize: {phrase!r}")
|
| 202 |
+
syl = syllable_count(phones)
|
| 203 |
+
div = len(set(phones)) / len(phones)
|
| 204 |
+
rar = rarity(phones)
|
| 205 |
+
conf, nearest = (None, None)
|
| 206 |
+
if not fast:
|
| 207 |
+
conf, nearest = confusability(phones)
|
| 208 |
+
|
| 209 |
+
# Syllable sweet spot: full credit for 3-5, penalized outside.
|
| 210 |
+
syl_score = {0: 0.0, 1: 0.2, 2: 0.6, 3: 1.0, 4: 1.0, 5: 1.0}.get(syl, 0.8)
|
| 211 |
+
|
| 212 |
+
# Acoustic reliability: mean over phones, with the onset (first phone)
|
| 213 |
+
# triple-weighted -- a swallowed first sound loses the whole match.
|
| 214 |
+
rel_vals = [RELIABILITY.get(p, 0.7) for p in phones]
|
| 215 |
+
reliability = (2 * rel_vals[0] + sum(rel_vals)) / (2 + len(rel_vals))
|
| 216 |
+
|
| 217 |
+
# Composite on a 0-100 scale. Weights favor confusability (the direct
|
| 218 |
+
# false-trigger proxy), rarity, and acoustic reliability.
|
| 219 |
+
parts = [
|
| 220 |
+
("rarity", min(rar / 14.0, 1.0), 0.25),
|
| 221 |
+
("confusability", min((conf or 0) / 0.45, 1.0), 0.30),
|
| 222 |
+
("reliability", reliability, 0.25),
|
| 223 |
+
("syllables", syl_score, 0.15),
|
| 224 |
+
("diversity", div, 0.05),
|
| 225 |
+
]
|
| 226 |
+
total = sum(v * w for _, v, w in parts) * 100
|
| 227 |
+
|
| 228 |
+
return {
|
| 229 |
+
"phrase": phrase,
|
| 230 |
+
"phones": " ".join(phones),
|
| 231 |
+
"syllables": syl,
|
| 232 |
+
"rarity_bits": round(rar, 2),
|
| 233 |
+
"min_distance": round(conf, 3) if conf is not None else None,
|
| 234 |
+
"nearest_common": nearest,
|
| 235 |
+
"reliability": round(reliability, 2),
|
| 236 |
+
"diversity": round(div, 2),
|
| 237 |
+
"score": round(total, 1),
|
| 238 |
+
}
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
# --- candidate generation ---------------------------------------------------
|
| 242 |
+
|
| 243 |
+
RARE_ONSETS = ["Z", "ZH", "V", "TH", "SH", "CH", "JH", "G", "K", "DR", "GR",
|
| 244 |
+
"KW", "SK", "SN", "PL", "KL", "TR", "FL"]
|
| 245 |
+
NUCLEI = ["AY", "OY", "AW", "EY", "OW", "IY", "UW", "AA", "ER"]
|
| 246 |
+
CODAS = ["", "K", "S", "KS", "N", "M", "SH", "Z", "NT", "RD", "L"]
|
| 247 |
+
|
| 248 |
+
# Respellings so TTS engines pronounce generated names predictably.
|
| 249 |
+
SPELL = {
|
| 250 |
+
"Z": "z", "ZH": "zh", "V": "v", "TH": "th", "SH": "sh", "CH": "ch",
|
| 251 |
+
"JH": "j", "G": "g", "K": "k", "DR": "dr", "GR": "gr", "KW": "qu",
|
| 252 |
+
"SK": "sk", "SN": "sn", "PL": "pl", "KL": "cl", "TR": "tr", "FL": "fl",
|
| 253 |
+
"AY": "y", "OY": "oy", "AW": "ow", "EY": "ay", "OW": "o", "IY": "ee",
|
| 254 |
+
"UW": "oo", "AA": "a", "ER": "er",
|
| 255 |
+
"S": "s", "KS": "x", "N": "n", "M": "m", "NT": "nt", "RD": "rd",
|
| 256 |
+
"L": "l", "": "",
|
| 257 |
+
}
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
def generate_candidates(n=200, seed=7):
|
| 261 |
+
rng = random.Random(seed)
|
| 262 |
+
out = set()
|
| 263 |
+
while len(out) < n:
|
| 264 |
+
syls = []
|
| 265 |
+
for _ in range(rng.choice([2, 3, 3])):
|
| 266 |
+
syls.append(rng.choice(RARE_ONSETS) + "|" + rng.choice(NUCLEI)
|
| 267 |
+
+ "|" + rng.choice(CODAS))
|
| 268 |
+
name = "".join(SPELL[p] for s in syls for p in s.split("|"))
|
| 269 |
+
if len(name) < 4:
|
| 270 |
+
continue
|
| 271 |
+
prefix = rng.choice(["hey ", "okay ", ""])
|
| 272 |
+
out.add(prefix + name)
|
| 273 |
+
return sorted(out)
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
def main():
|
| 277 |
+
ap = argparse.ArgumentParser()
|
| 278 |
+
ap.add_argument("phrase", nargs="?", help="phrase to score")
|
| 279 |
+
ap.add_argument("--rank", help="file with one phrase per line")
|
| 280 |
+
ap.add_argument("--suggest", type=int, metavar="N",
|
| 281 |
+
help="generate candidates and print top N")
|
| 282 |
+
ap.add_argument("--json", action="store_true")
|
| 283 |
+
args = ap.parse_args()
|
| 284 |
+
|
| 285 |
+
if args.phrase:
|
| 286 |
+
result = score_phrase(args.phrase)
|
| 287 |
+
print(json.dumps(result, indent=2) if args.json else format_row(result, header=True))
|
| 288 |
+
return
|
| 289 |
+
|
| 290 |
+
phrases = []
|
| 291 |
+
if args.rank:
|
| 292 |
+
with open(args.rank, encoding="utf-8") as f:
|
| 293 |
+
phrases = [l.strip() for l in f if l.strip()]
|
| 294 |
+
elif args.suggest:
|
| 295 |
+
print("generating and pre-screening candidates...", file=sys.stderr)
|
| 296 |
+
cands = generate_candidates()
|
| 297 |
+
# cheap pass on rarity/syllables first, full scoring on survivors
|
| 298 |
+
pre = []
|
| 299 |
+
for c in cands:
|
| 300 |
+
try:
|
| 301 |
+
ph = phones_for_phrase(c)
|
| 302 |
+
except Exception:
|
| 303 |
+
continue
|
| 304 |
+
if 3 <= syllable_count(ph) <= 5:
|
| 305 |
+
pre.append((rarity(ph), c))
|
| 306 |
+
pre.sort(reverse=True)
|
| 307 |
+
phrases = [c for _, c in pre[: args.suggest * 3]]
|
| 308 |
+
else:
|
| 309 |
+
ap.print_help()
|
| 310 |
+
return
|
| 311 |
+
|
| 312 |
+
results = []
|
| 313 |
+
for p in phrases:
|
| 314 |
+
try:
|
| 315 |
+
results.append(score_phrase(p))
|
| 316 |
+
except ValueError as e:
|
| 317 |
+
print(f"skipped: {e}", file=sys.stderr)
|
| 318 |
+
results.sort(key=lambda r: -r["score"])
|
| 319 |
+
if args.suggest:
|
| 320 |
+
results = results[: args.suggest]
|
| 321 |
+
if args.json:
|
| 322 |
+
print(json.dumps(results, indent=2))
|
| 323 |
+
else:
|
| 324 |
+
for i, r in enumerate(results):
|
| 325 |
+
print(format_row(r, header=(i == 0)))
|
| 326 |
+
|
| 327 |
+
|
| 328 |
+
def format_row(r, header=False):
|
| 329 |
+
row = (f"{r['score']:6.1f} {r['phrase']:<20} syl={r['syllables']} "
|
| 330 |
+
f"rel={r['reliability']:.2f} rarity={r['rarity_bits']:5.2f}b "
|
| 331 |
+
f"dist={r['min_distance']} (nearest: {r['nearest_common']}) "
|
| 332 |
+
f"[{r['phones']}]")
|
| 333 |
+
if header:
|
| 334 |
+
return (" score phrase details\n"
|
| 335 |
+
" ----- ------ -------\n" + row)
|
| 336 |
+
return row
|
| 337 |
+
|
| 338 |
+
|
| 339 |
+
if __name__ == "__main__":
|
| 340 |
+
main()
|
phoneme_engine/spot_file.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Offline spotting: run a wake word over a wav file (or a directory).
|
| 2 |
+
|
| 3 |
+
python -m phoneme_engine.spot_file "hey orbit" clip.wav --threshold -2.5
|
| 4 |
+
|
| 5 |
+
Prints detections with timestamps, plus the greedy phoneme transcription
|
| 6 |
+
with --show-phones. Used for automated testing (TTS-generated positives,
|
| 7 |
+
LibriSpeech negatives) without a microphone.
|
| 8 |
+
"""
|
| 9 |
+
|
| 10 |
+
import argparse
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
|
| 13 |
+
import soundfile as sf
|
| 14 |
+
import torch
|
| 15 |
+
import torchaudio
|
| 16 |
+
|
| 17 |
+
from .decoder import KeywordSpotter
|
| 18 |
+
from .features import SAMPLE_RATE, LogMel
|
| 19 |
+
from .model import PhonemeTCN
|
| 20 |
+
from .phones import ids_to_phones
|
| 21 |
+
|
| 22 |
+
ROOT = Path(__file__).resolve().parent.parent
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def load_wav(path, device):
|
| 26 |
+
audio, sr = sf.read(str(path), dtype="float32")
|
| 27 |
+
wav = torch.from_numpy(audio)
|
| 28 |
+
if wav.ndim == 2:
|
| 29 |
+
wav = wav.mean(dim=1)
|
| 30 |
+
wav = wav.unsqueeze(0)
|
| 31 |
+
if sr != SAMPLE_RATE:
|
| 32 |
+
wav = torchaudio.functional.resample(wav, sr, SAMPLE_RATE)
|
| 33 |
+
return wav.to(device)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def main():
|
| 37 |
+
ap = argparse.ArgumentParser()
|
| 38 |
+
ap.add_argument("phrase")
|
| 39 |
+
ap.add_argument("path")
|
| 40 |
+
ap.add_argument("--threshold", type=float, default=-4.4)
|
| 41 |
+
ap.add_argument("--ckpt", default=str(ROOT / "checkpoints" / "best.pt"))
|
| 42 |
+
ap.add_argument("--show-phones", action="store_true")
|
| 43 |
+
args = ap.parse_args()
|
| 44 |
+
|
| 45 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 46 |
+
frontend = LogMel().to(device).eval()
|
| 47 |
+
model = PhonemeTCN().to(device).eval()
|
| 48 |
+
state = torch.load(args.ckpt, map_location=device, weights_only=True)
|
| 49 |
+
model.load_state_dict(state["model"])
|
| 50 |
+
|
| 51 |
+
p = Path(args.path)
|
| 52 |
+
files = sorted(p.glob("**/*.wav")) + sorted(p.glob("**/*.flac")) \
|
| 53 |
+
if p.is_dir() else [p]
|
| 54 |
+
|
| 55 |
+
total_hits = 0
|
| 56 |
+
with torch.no_grad():
|
| 57 |
+
for f in files:
|
| 58 |
+
wav = load_wav(f, device)
|
| 59 |
+
logits = model(frontend(wav))
|
| 60 |
+
lp = torch.log_softmax(logits.float(), dim=2)[0].cpu().numpy()
|
| 61 |
+
if args.show_phones:
|
| 62 |
+
ids = lp.argmax(axis=1).tolist()
|
| 63 |
+
print(f"{f.name} phones: {' '.join(ids_to_phones(ids))}")
|
| 64 |
+
spotter = KeywordSpotter(args.phrase, threshold=args.threshold)
|
| 65 |
+
hits = spotter.run(lp)
|
| 66 |
+
total_hits += len(hits)
|
| 67 |
+
for frame, score in hits:
|
| 68 |
+
print(f"{f.name}: DETECT at {frame * 0.02:.2f}s "
|
| 69 |
+
f"score={score:.2f}")
|
| 70 |
+
if not hits:
|
| 71 |
+
print(f"{f.name}: no detection")
|
| 72 |
+
print(f"total detections: {total_hits} across {len(files)} file(s)")
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
if __name__ == "__main__":
|
| 76 |
+
main()
|
phoneme_tcn_student.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:78477bd4175f161df058b48ff084cfd01ab05c1f980902bc9c8f62d32cba03f5
|
| 3 |
+
size 4410358
|
phoneme_tcn_teacher.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:dae5cdd457d97d4297c0392389622ca348202c1a451f5c04e79bc007d81b6f10
|
| 3 |
+
size 40277054
|