JoelAjitesh commited on
Commit
f6aec75
·
verified ·
1 Parent(s): d3353c9

Phoneme wake word engine: student+teacher models, INT8 export, C engine, enrollment tooling

Browse files
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