mirror of
https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git
synced 2026-08-12 16:05:34 -06:00
Compare commits
46 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c474deb8b5 | ||
|
|
931694b711 | ||
|
|
5554b2eb5e | ||
|
|
7d77f71dc3 | ||
|
|
3d341d0617 | ||
|
|
a1b22200e0 | ||
|
|
89260f1f14 | ||
|
|
0140dfb56f | ||
|
|
1fc7d80bae | ||
|
|
31a6388da4 | ||
|
|
85c2d6334b | ||
|
|
5f6f108c85 | ||
|
|
bb5033c5fb | ||
|
|
8a8f4a82d9 | ||
|
|
ed120e91ab | ||
|
|
7d8ebd6637 | ||
|
|
874f273d0b | ||
|
|
04249f414d | ||
|
|
6a0d60d569 | ||
|
|
8df17599c2 | ||
|
|
280e8f8de4 | ||
|
|
b582a6cade | ||
|
|
196ab8c0e7 | ||
|
|
134f607bef | ||
|
|
4a9e2f2cde | ||
|
|
7c246856df | ||
|
|
3705dabc09 | ||
|
|
1dcf48209f | ||
|
|
4f44bef8d5 | ||
|
|
98fa879db1 | ||
|
|
dfac549430 | ||
|
|
775a78326b | ||
|
|
429be4cc67 | ||
|
|
2e6179ec32 | ||
|
|
ee6ff6e9d5 | ||
|
|
251b0280b6 | ||
|
|
dd2bdda431 | ||
|
|
9d8e0afe1b | ||
|
|
318a4ad3b5 | ||
|
|
51cbf6fd90 | ||
|
|
6e7396455a | ||
|
|
2da9f7a686 | ||
|
|
7b028e4420 | ||
|
|
b3d9f0e369 | ||
|
|
240ca7682e | ||
|
|
18e5fcd000 |
144
.github/workflows/docker-publish.yml
vendored
Normal file
144
.github/workflows/docker-publish.yml
vendored
Normal file
@@ -0,0 +1,144 @@
|
||||
name: Publish Docker Images
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- "v*"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
packages: write
|
||||
|
||||
concurrency:
|
||||
group: docker-publish-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
env:
|
||||
REGISTRY: ghcr.io
|
||||
IMAGE_NAME: tatertotterson/microwakeword
|
||||
|
||||
jobs:
|
||||
docker:
|
||||
name: Docker image
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Validate tag matches trainer version
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
shell: bash
|
||||
run: |
|
||||
set -euo pipefail
|
||||
version="$(tr -d '[:space:]' < VERSION)"
|
||||
expected_tag="v${version#v}"
|
||||
if [[ "${GITHUB_REF_NAME}" != "${expected_tag}" ]]; then
|
||||
echo "Tag ${GITHUB_REF_NAME} does not match trainer version ${expected_tag}." >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to GHCR
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ${{ env.REGISTRY }}
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Docker metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}
|
||||
flavor: |
|
||||
latest=false
|
||||
tags: |
|
||||
type=raw,value=latest
|
||||
type=ref,event=tag
|
||||
|
||||
- name: Build and push image
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: dockerfile
|
||||
platforms: linux/amd64
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha,scope=mww-trainer-nvidia-docker
|
||||
cache-to: type=gha,mode=max,scope=mww-trainer-nvidia-docker
|
||||
|
||||
- name: Docker metadata (Blackwell)
|
||||
id: meta-blackwell
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}
|
||||
flavor: |
|
||||
latest=false
|
||||
tags: |
|
||||
type=raw,value=blackwell
|
||||
type=ref,event=tag,suffix=-blackwell
|
||||
|
||||
- name: Build and push Blackwell image
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: dockerfile.blackwell
|
||||
platforms: linux/amd64
|
||||
push: true
|
||||
tags: ${{ steps.meta-blackwell.outputs.tags }}
|
||||
labels: ${{ steps.meta-blackwell.outputs.labels }}
|
||||
cache-from: type=gha,scope=mww-trainer-nvidia-docker-blackwell
|
||||
cache-to: type=gha,mode=max,scope=mww-trainer-nvidia-docker-blackwell
|
||||
|
||||
- name: Create release notes
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
TAG_NAME: ${{ github.ref_name }}
|
||||
REPO: ${{ github.repository }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
|
||||
title="microWakeWord Nvidia Trainer ${TAG_NAME}"
|
||||
generated_notes="$(mktemp)"
|
||||
release_notes="$(mktemp)"
|
||||
test -s WHATS_NEW.md
|
||||
|
||||
gh api "repos/${REPO}/releases/generate-notes" \
|
||||
-f tag_name="${TAG_NAME}" \
|
||||
-f target_commitish="${GITHUB_SHA}" \
|
||||
--jq '.body' > "${generated_notes}"
|
||||
|
||||
{
|
||||
echo "## What's New"
|
||||
echo
|
||||
cat WHATS_NEW.md
|
||||
echo
|
||||
echo
|
||||
echo "## Docker Images"
|
||||
echo
|
||||
echo "- \`ghcr.io/tatertotterson/microwakeword:${TAG_NAME}\`"
|
||||
echo "- \`ghcr.io/tatertotterson/microwakeword:latest\`"
|
||||
echo "- \`ghcr.io/tatertotterson/microwakeword:${TAG_NAME}-blackwell\`"
|
||||
echo "- \`ghcr.io/tatertotterson/microwakeword:blackwell\`"
|
||||
echo
|
||||
cat "${generated_notes}"
|
||||
} > "${release_notes}"
|
||||
|
||||
if gh release view "${TAG_NAME}" >/dev/null 2>&1; then
|
||||
gh release edit "${TAG_NAME}" \
|
||||
--title "${title}" \
|
||||
--notes-file "${release_notes}" \
|
||||
--latest \
|
||||
--verify-tag
|
||||
else
|
||||
gh release create "${TAG_NAME}" \
|
||||
--title "${title}" \
|
||||
--notes-file "${release_notes}" \
|
||||
--latest \
|
||||
--verify-tag
|
||||
fi
|
||||
1
.gitignore
vendored
1
.gitignore
vendored
@@ -1,3 +1,4 @@
|
||||
personal_samples/*
|
||||
data/
|
||||
trim_history/
|
||||
.DS_Store
|
||||
423
README.md
423
README.md
@@ -1,169 +1,346 @@
|
||||
<div align="center">
|
||||
<h1>🎙️ microWakeWord Nvidia Trainer & Recorder</h1>
|
||||
<img width="1002" height="593" alt="Screenshot 2026-01-18 at 8 13 35 AM" src="https://github.com/user-attachments/assets/e1411d8a-8638-4df8-992b-09a46c6e5ddc" />
|
||||
<a href="https://taterassistant.com">
|
||||
<img src="images/tater-repo-logo.png" alt="microWakeWord Trainer" width="460"/>
|
||||
</a>
|
||||
</div>
|
||||
<h3 align="center">
|
||||
<a href="https://taterassistant.com">taterassistant.com</a>
|
||||
</h3>
|
||||
|
||||
Train **microWakeWord** detection models using a simple **web-based recorder + trainer UI**, packaged in a Docker container.
|
||||
Train custom microWakeWord models in Docker with NVIDIA/CUDA acceleration, generated Piper samples, device-captured samples, reviewed false-wake negatives, live training logs, and local wake-word links for Tater Native satellites.
|
||||
|
||||
No Jupyter notebooks required. No manual cell execution. Just record your voice (optional) and train.
|
||||
Real samples come from device-captured wake audio, close misses, or manual uploads. Every saved sample is normalized to `16 kHz / mono / 16-bit PCM WAV` before training.
|
||||
|
||||
---
|
||||
|
||||
<img width="100" height="44" alt="unraid_logo_black-339076895" src="https://github.com/user-attachments/assets/87351bed-3321-4a43-924f-fecf2e4e700f" />
|
||||
|
||||
**microWakeWord_Trainer-Nvidia** is available in the **Unraid Community Apps** store.
|
||||
Install directly from the Unraid App Store with a one-click template.
|
||||
|
||||
---
|
||||
|
||||
<img width="100" height="56" alt="unraid_logo_black-339076895" src="https://github.com/user-attachments/assets/bf959585-ae13-4b4d-ae62-4202a850d35a" />
|
||||
|
||||
|
||||
### Pull the Docker Image
|
||||
## Docker Image
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:latest
|
||||
```
|
||||
|
||||
Tagged releases also publish matching immutable image tags:
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v15
|
||||
```
|
||||
|
||||
The release tag must match `VERSION`. Update `WHATS_NEW.md` before tagging; the Docker workflow prepends it to GitHub's automatically generated release notes.
|
||||
|
||||
RTX 50-series / Blackwell GPUs use a separate image with CUDA 12.8 and a
|
||||
Python 3.13 TensorFlow build for `sm_120`:
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:blackwell
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v15-blackwell
|
||||
```
|
||||
|
||||
Use the Blackwell image only for RTX 50-series cards. It includes the
|
||||
community-built TensorFlow wheel from
|
||||
[chivitiH/tensorflow-blackwell-python313](https://github.com/chivitiH/tensorflow-blackwell-python313),
|
||||
which is unofficial and licensed CC BY-NC 4.0.
|
||||
|
||||
---
|
||||
|
||||
### Run the Container
|
||||
## Run The Container
|
||||
|
||||
```bash
|
||||
docker run -d \
|
||||
--gpus all \
|
||||
-p 8888:8888 \
|
||||
--network host \
|
||||
-e REC_PORT=8789 \
|
||||
-v $(pwd):/data \
|
||||
ghcr.io/tatertotterson/microwakeword:latest
|
||||
```
|
||||
|
||||
**What these flags do:**
|
||||
- `--gpus all` → Enables GPU acceleration
|
||||
- `-p 8888:8888` → Exposes the Recorder + Trainer WebUI
|
||||
- `-v $(pwd):/data` → Persists all models, datasets, and cache
|
||||
Use a version tag such as `ghcr.io/tatertotterson/microwakeword:v15` when you want to pin a known release instead of tracking `latest`.
|
||||
For RTX 50-series cards, use `ghcr.io/tatertotterson/microwakeword:blackwell`
|
||||
or a pinned tag such as `ghcr.io/tatertotterson/microwakeword:v15-blackwell`
|
||||
in the same `docker run` command.
|
||||
|
||||
---
|
||||
The flags:
|
||||
|
||||
### Open the Recorder WebUI
|
||||
- `--gpus all` enables GPU acceleration.
|
||||
- `--network host` exposes the trainer server directly so satellites can send captured audio and load trained wake-word files.
|
||||
- `-e REC_PORT=8789` sets the trainer web UI and captured-audio port. Change this value if `8789` is already in use.
|
||||
- `-v $(pwd):/data` persists models, downloaded voices, datasets, samples, and generated wake-word artifacts.
|
||||
|
||||
Open your browser and go to:
|
||||
If you do not use host networking, publish the trainer port and make sure satellites can reach it from your LAN.
|
||||
|
||||
👉 **http://localhost:8888**
|
||||
|
||||
You’ll see the **microWakeWord Recorder & Trainer UI**.
|
||||
|
||||
---
|
||||
|
||||
## 🎤 Recording Voice Samples (Optional)
|
||||
|
||||
Personal voice recordings are **optional**.
|
||||
|
||||
- You may **record your own voice** for better accuracy
|
||||
- Or simply **click “Train” without recording anything**
|
||||
|
||||
If no recordings are present, training will proceed using **synthetic TTS samples only**.
|
||||
|
||||
### Remote systems (important)
|
||||
If you are running this on a **remote PC / server**, browser-based recording will not work unless:
|
||||
- You use a **reverse proxy** (HTTPS + mic permissions), **or**
|
||||
- You access the UI via **localhost** on the same machine
|
||||
|
||||
Training itself works fine remotely — only recording requires local microphone access.
|
||||
|
||||
---
|
||||
|
||||
### 🎙️ Recording Flow
|
||||
|
||||
1. Enter your wake word
|
||||
2. Test pronunciation with **Test TTS**
|
||||
3. Choose:
|
||||
- Number of speakers (e.g. family members)
|
||||
- Takes per speaker (default: 10)
|
||||
4. Click **Begin recording**
|
||||
5. Speak naturally — recording:
|
||||
- Starts when you talk
|
||||
- Stops automatically after silence
|
||||
6. Repeat for each speaker
|
||||
|
||||
Files are saved automatically to:
|
||||
|
||||
```
|
||||
personal_samples/
|
||||
speaker01_take01.wav
|
||||
speaker01_take02.wav
|
||||
speaker02_take01.wav
|
||||
...
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 🧠 Training Behavior (Important Notes)
|
||||
|
||||
### ⏬ First training run
|
||||
The **first time you click Train**, the system will download **large training datasets** (background noise, speech corpora, etc.).
|
||||
|
||||
- This can take **several minutes**
|
||||
- This happens **only once**
|
||||
- Data is cached inside `/data`
|
||||
|
||||
You **will NOT need to download these again** unless you delete `/data`.
|
||||
|
||||
---
|
||||
|
||||
### 🔁 Re-training is safe and incremental
|
||||
|
||||
- You can train **multiple wake words** back-to-back
|
||||
- You do **NOT** need to clear any folders between runs
|
||||
- Old models are preserved in timestamped output directories
|
||||
- All required cleanup and reuse logic is handled automatically
|
||||
|
||||
---
|
||||
|
||||
## 📦 Output Files
|
||||
|
||||
When training completes, you’ll get:
|
||||
- `<wake_word>.tflite` – quantized streaming model
|
||||
- `<wake_word>.json` – ESPHome-compatible metadata
|
||||
|
||||
Both are saved under:
|
||||
Open:
|
||||
|
||||
```text
|
||||
/data/output/
|
||||
http://localhost:8789
|
||||
```
|
||||
|
||||
Each run is placed in its own timestamped folder.
|
||||
If you change `REC_PORT`, open that port instead and use the same port in the satellite `Trainer App URL`.
|
||||
|
||||
---
|
||||
|
||||
## 🎤 Optional: Personal Voice Samples (Advanced)
|
||||
## What The UI Does
|
||||
|
||||
If you record personal samples:
|
||||
- They are automatically augmented
|
||||
- They are **up-weighted during training**
|
||||
- This significantly improves real-world accuracy
|
||||
|
||||
No configuration required — detection is automatic.
|
||||
- `Trainer` starts a wake-word session, shows positive/negative sample counts, and launches training.
|
||||
- `Auto Training` transcribes real wake triggers, promotes phrase-misses to hard negatives, schedules retraining, and refreshes Tater Native satellites.
|
||||
- `Captured Audio` reviews clips sent by Tater Native or ESPHome sats, including wake hits, close misses, and false wakes.
|
||||
- `Samples` plays, removes, clears, and manually imports personal or negative samples.
|
||||
- `Wake Words` lists locally trained JSON/model links for live wake-word switching in Tater.
|
||||
- Popup consoles show colorized training logs while long-running jobs are active.
|
||||
|
||||
---
|
||||
|
||||
## 🔄 Resetting Everything (Optional)
|
||||
## Captured Audio Workflow
|
||||
|
||||
If you want a **completely clean slate**:
|
||||
To collect samples from a sat, point its trainer feedback setting at this app. Tater Native satellites use the native settings popup in Tater. Older ESPHome satellites can still use their device entities.
|
||||
|
||||
Delete the /data folder
|
||||
For Tater Native satellites, enable trainer feedback in Tater:
|
||||
|
||||
Then restart the container.
|
||||
- `Send Good Wakes To Trainer` toggles upload of confirmed wake-word triggers.
|
||||
- `Send Close Misses To Trainer` toggles upload of near misses.
|
||||
- `Trainer App URL` sets the trainer address, for example `http://trainer.local:8789` or `http://<trainer-ip>:8789`.
|
||||
|
||||
⚠️ This will:
|
||||
- Remove cached datasets
|
||||
- Require re-downloading training data
|
||||
- Delete trained models
|
||||
For older ESPHome firmware, the equivalent capture setup is exposed as device entities:
|
||||
|
||||
- `Capture Wake Audio` toggles upload of wake-word triggers.
|
||||
- `Capture Close Misses` toggles upload of near misses.
|
||||
- `Trainer App URL` sets the trainer address, for example `http://<trainer-ip>:8789`.
|
||||
|
||||
Satellites send raw captured audio to:
|
||||
|
||||
```text
|
||||
/api/upload_captured_audio_raw
|
||||
```
|
||||
|
||||
Keep the training app running and reachable at the `Trainer App URL` while capture is enabled. The sats upload clips live; if the app is stopped or the URL is wrong, captured audio will not be saved.
|
||||
|
||||
In the `Captured Audio` tab:
|
||||
|
||||
- play each clip from the inbox
|
||||
- mark good wake-word clips as `This is good`
|
||||
- mark bad triggers as `False wake`
|
||||
- discard clips that should not be used
|
||||
|
||||
Approved clips move into:
|
||||
|
||||
```text
|
||||
/data/personal_samples/
|
||||
```
|
||||
|
||||
False wakes move into:
|
||||
|
||||
```text
|
||||
/data/negative_samples/
|
||||
```
|
||||
|
||||
Captured audio is boosted for easier playback in the UI, then kept in the correct training format.
|
||||
|
||||
---
|
||||
|
||||
## 🙌 Credits
|
||||
## Samples
|
||||
|
||||
Built on top of the excellent
|
||||
**https://github.com/kahrendt/microWakeWord**
|
||||
The `Samples` tab is the sample library.
|
||||
|
||||
Huge thanks to the original authors ❤️
|
||||
- `Personal` samples are positive examples of the wake word.
|
||||
- `Negative` samples are reviewed false wakes or hard negatives.
|
||||
- Both can be played back and removed one at a time.
|
||||
- Manual upload is available here as an optional seed path.
|
||||
|
||||
Accepted manual upload formats include:
|
||||
|
||||
- WAV
|
||||
- MP3
|
||||
- M4A
|
||||
- FLAC
|
||||
- OGG
|
||||
- AAC
|
||||
- OPUS
|
||||
- WEBM
|
||||
|
||||
Uploads are validated or converted with `ffmpeg` into:
|
||||
|
||||
```text
|
||||
16 kHz / mono / 16-bit PCM WAV
|
||||
```
|
||||
|
||||
Starting a new session does not clear samples. Use the clear buttons in `Samples` if you want to remove saved personal or negative clips.
|
||||
|
||||
---
|
||||
|
||||
## Auto Training
|
||||
|
||||
`Auto Training` is an opt-in sample-review and retraining loop. It is disabled until you enter the exact wake phrase and enable it.
|
||||
|
||||
For each new wake-trigger clip sent to the trainer:
|
||||
|
||||
1. Faster Whisper transcribes the audio locally.
|
||||
2. If the transcript contains the configured wake phrase, the clip stays in `Captured Audio` for manual review by default.
|
||||
3. If speech was transcribed but the wake phrase is absent, the clip moves to `/data/negative_samples/` as an auto-reviewed hard negative.
|
||||
4. Empty transcripts, VAD-blocked captures, and captures for another wake word stay out of the automatic negative path.
|
||||
|
||||
Two optional cleanup rules are available:
|
||||
|
||||
- `Delete confirmed good wakes` removes normal wake-trigger clips after STT confirms the configured phrase.
|
||||
- `Promote confirmed close misses` checks close misses that passed VAD and moves them to the personal positive samples only when STT confirms the configured phrase.
|
||||
|
||||
A close miss with an empty transcript or without the configured phrase stays in `Captured Audio`; it is never turned into a negative automatically. Saving Auto Training settings also scans existing eligible captures. Enabling close-miss promotion reviews previous unreviewed close misses, while enabling cleanup removes previously confirmed good wakes without transcribing them a second time.
|
||||
|
||||
The default `small.en` model uses CUDA with `float16` when CTranslate2 can see an NVIDIA GPU, and falls back to CPU with `int8`. Choose a multilingual Faster Whisper model such as `small` when the wake phrase is not English. Downloaded STT models are cached in `/data/auto_train_models/`.
|
||||
|
||||
Scheduled training runs only after the configured number of new automatic negatives has accumulated. A successful run securely publishes the trained wake-word name and JSON URL to the linked Tater instance. Tater saves it as the global satellite wake word and pushes the updated setting to every connected satellite, so no satellite firmware change is required.
|
||||
|
||||
The `Trainer public URL` must be reachable from the satellites. With the documented `--network host` command, the trainer can normally use the LAN address from the browser request or host network. If you open the UI as `http://localhost:8789`, enter a value such as `http://192.168.1.50:8789`, or start the container with `REC_PUBLIC_BASE_URL` set to that value. When using Docker bridge networking, always set this URL to the published host address; a container bridge address is not satellite-reachable.
|
||||
|
||||
The default Tater URL, `http://127.0.0.1:8501`, assumes the documented host networking. Change it to a container-reachable Tater address if you use another Docker network. Click `Link Tater` and enter the short-lived code shown in Tater Voice Settings; the resulting trainer-specific link credential is stored in `/data/auto_train_config.json` with owner-only permissions.
|
||||
|
||||
---
|
||||
|
||||
## Training Flow
|
||||
|
||||
1. Enter the wake phrase in `Trainer`.
|
||||
2. Choose the language.
|
||||
3. Optionally test pronunciation with `Test TTS`.
|
||||
4. Review the positive and negative sample counts.
|
||||
5. Click `Start training`.
|
||||
6. Watch the popup training console.
|
||||
|
||||
Personal samples are optional. Training can run with zero personal samples after confirmation, using generated TTS samples and the stock negative datasets.
|
||||
|
||||
Reviewed negative samples are converted into `/data/work/reviewed_negative_features/` and inserted into the training YAML as a hard-negative feature set when present.
|
||||
|
||||
On RTX 50-series / Blackwell GPUs, the Blackwell Docker image keeps sample generation and augmentation in the normal Python 3.12 trainer environment, then runs only the TensorFlow training/export stage in `/data/.venv-blackwell` with Python 3.13 and the Blackwell-native TensorFlow wheel.
|
||||
|
||||
---
|
||||
|
||||
## Language Support
|
||||
|
||||
The language picker is dynamic.
|
||||
|
||||
- `en` is always available.
|
||||
- English keeps the existing dedicated generator model path.
|
||||
- Non-English languages are discovered from the Piper voices catalog and any local Piper voice metadata.
|
||||
- When a non-English language is selected, the trainer downloads all voices for that selected language only.
|
||||
- Already-downloaded voices are reused.
|
||||
- It does not download every language up front.
|
||||
|
||||
If the upstream Piper catalog is unavailable, already-installed local voices are used when available.
|
||||
|
||||
---
|
||||
|
||||
## Dataset Behavior
|
||||
|
||||
The first training run downloads and prepares missing training assets into `/data`, including:
|
||||
|
||||
- Piper voices for the selected language
|
||||
- negative datasets and background data
|
||||
- the Python training environment
|
||||
- generated samples and augmented feature caches
|
||||
|
||||
After those assets are prepared, later runs reuse the local copies unless the mounted `/data` contents are deleted.
|
||||
|
||||
---
|
||||
|
||||
## Trained Wake Words
|
||||
|
||||
The `Wake Words` tab lists locally trained wake-word packages from `/data/trained_wake_words/`.
|
||||
|
||||
- Copy the JSON URL into the Tater Native satellite settings to switch wake words live.
|
||||
- Links use the configured public trainer URL, a non-loopback browser host, or the detected LAN address instead of advertising `127.0.0.1` to satellites.
|
||||
- Open the JSON or model links directly for quick inspection.
|
||||
- The JSON includes the matching model path plus Tater tuning metadata.
|
||||
- No firmware flashing happens from this trainer app anymore.
|
||||
|
||||
Use the main Tater app for satellite firmware updates and USB flashing.
|
||||
|
||||
---
|
||||
|
||||
## Output Files
|
||||
|
||||
Successful runs produce timestamped training output folders such as:
|
||||
|
||||
```text
|
||||
/data/output/<timestamp>-<wake_word>-<samples>-<steps>/<wake_word>.tflite
|
||||
/data/output/<timestamp>-<wake_word>-<samples>-<steps>/<wake_word>.json
|
||||
```
|
||||
|
||||
The trainer also syncs Tater-ready wake-word artifacts into:
|
||||
|
||||
```text
|
||||
/data/trained_wake_words/<wake_word>.tflite
|
||||
/data/trained_wake_words/<wake_word>.json
|
||||
```
|
||||
|
||||
The `Wake Words` tab uses `/data/trained_wake_words/` to populate the local wake-word links.
|
||||
|
||||
The JSON keeps the standard microWakeWord fields for compatibility:
|
||||
|
||||
```json
|
||||
{
|
||||
"micro": {
|
||||
"probability_cutoff": 0.97,
|
||||
"sliding_window_size": 6
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
It also includes Tater Native metadata used by newer satellites and the Tater settings UI:
|
||||
|
||||
```json
|
||||
{
|
||||
"model_format": "tflite_stream_state_internal_quant",
|
||||
"quantization": "int8",
|
||||
"sample_rate": 16000,
|
||||
"tater_native": {
|
||||
"format_version": 1,
|
||||
"wake_threshold": 0.97,
|
||||
"wake_sliding_window": 6,
|
||||
"close_miss_threshold": 0.80,
|
||||
"frontend": {
|
||||
"name": "tflm_microfrontend",
|
||||
"sample_rate": 16000,
|
||||
"feature_duration_ms": 30,
|
||||
"feature_step_ms": 10,
|
||||
"feature_size": 40
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Calibration metrics are included under `calibration` so false accepts/hour and recall can be surfaced in the UI.
|
||||
Calibration evaluates thresholds from `0.95` through `1.00` with sliding windows of `5`, `6`, and `7`. Among candidates within 0.5 percentage points of the best recall, it prefers the lowest measured ambient false-accept rate. If calibration cannot complete, packaging uses the conservative `0.97` threshold and a window of `6`.
|
||||
|
||||
---
|
||||
|
||||
## Resetting Everything
|
||||
|
||||
If you want a clean slate, stop the container and remove the contents of the mounted `/data` directory.
|
||||
|
||||
That removes:
|
||||
|
||||
- personal samples
|
||||
- negative samples
|
||||
- captured inbox clips
|
||||
- downloaded Piper voices
|
||||
- cached datasets
|
||||
- training environments
|
||||
- trained models
|
||||
- Auto Training settings, state, transcripts, and cached Faster Whisper models
|
||||
|
||||
---
|
||||
|
||||
## Important Notes
|
||||
|
||||
- Personal samples are optional.
|
||||
- Negative samples are optional but useful for reducing false wakes.
|
||||
- Auto Training is disabled by default and only classifies actual wake triggers automatically.
|
||||
- The UI server is `trainer_server.py`.
|
||||
- The launcher is `run.sh`.
|
||||
- Trainer capture settings live in Tater for Tater Native satellites, and on device entities for older ESPHome satellites.
|
||||
|
||||
---
|
||||
|
||||
## Credits
|
||||
|
||||
Built on top of:
|
||||
|
||||
- [microWakeWord](https://github.com/kahrendt/microWakeWord)
|
||||
- [piper-sample-generator](https://github.com/rhasspy/piper-sample-generator)
|
||||
- [tensorflow-blackwell-python313](https://github.com/chivitiH/tensorflow-blackwell-python313) for the optional RTX 50-series / Blackwell image
|
||||
|
||||
3
WHATS_NEW.md
Normal file
3
WHATS_NEW.md
Normal file
@@ -0,0 +1,3 @@
|
||||
- Added secure Tater linking: enter the short-lived code from Tater Voice Settings instead of giving the trainer a general API token.
|
||||
- Automatic and manual publishing now tell Tater which trained wake word is active, and Tater applies it globally to every connected satellite.
|
||||
- Added clear linked, unlinked, and pairing-success states to the Auto Training interface.
|
||||
451
cli/calibrate_detector.py
Normal file
451
cli/calibrate_detector.py
Normal file
@@ -0,0 +1,451 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Choose detector metadata that better matches the trained model."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Sequence
|
||||
|
||||
import numpy as np
|
||||
import yaml
|
||||
|
||||
DEFAULT_WINDOW_SIZES = [5, 6, 7]
|
||||
DEFAULT_TARGET_FAPH = float(os.environ.get("MWW_CALIBRATION_TARGET_FAPH", "0.25"))
|
||||
DEFAULT_COOLDOWN_SLICES = int(os.environ.get("MWW_CALIBRATION_COOLDOWN_SLICES", "25"))
|
||||
DEFAULT_POSITIVE_SKIP_SLICES = int(
|
||||
os.environ.get("MWW_CALIBRATION_POSITIVE_SKIP_SLICES", "25")
|
||||
)
|
||||
DEFAULT_CUTOFF_STEP = float(os.environ.get("MWW_CALIBRATION_CUTOFF_STEP", "0.01"))
|
||||
DEFAULT_CUTOFF_MIN = float(os.environ.get("MWW_CALIBRATION_CUTOFF_MIN", "0.95"))
|
||||
DEFAULT_CUTOFF_MAX = float(os.environ.get("MWW_CALIBRATION_CUTOFF_MAX", "1.00"))
|
||||
DEFAULT_RECALL_MARGIN = float(os.environ.get("MWW_CALIBRATION_RECALL_MARGIN", "0.005"))
|
||||
PREFERRED_WINDOW_SIZE = 6
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Calibrate microWakeWord detector metadata from validation data."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--training-config",
|
||||
default="trained_models/wakeword/training_config.yaml",
|
||||
help="Path to the saved microWakeWord training_config.yaml file.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default=(
|
||||
"trained_models/wakeword/tflite_stream_state_internal_quant/"
|
||||
"stream_state_internal_quant.tflite"
|
||||
),
|
||||
help="Path to the quantized streaming TFLite model.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default=(
|
||||
"trained_models/wakeword/tflite_stream_state_internal_quant/"
|
||||
"detection_calibration.json"
|
||||
),
|
||||
help="Where to write the selected detector settings as JSON.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--window-sizes",
|
||||
default=",".join(str(value) for value in DEFAULT_WINDOW_SIZES),
|
||||
help="Comma-separated sliding window sizes to evaluate.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--target-faph",
|
||||
type=float,
|
||||
default=DEFAULT_TARGET_FAPH,
|
||||
help="Target ambient false accepts per hour for the selected operating point.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--recall-margin",
|
||||
type=float,
|
||||
default=DEFAULT_RECALL_MARGIN,
|
||||
help=(
|
||||
"Maximum recall loss allowed when preferring a candidate with fewer "
|
||||
"ambient false accepts (0.005 means 0.5 percentage points)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cooldown-slices",
|
||||
type=int,
|
||||
default=DEFAULT_COOLDOWN_SLICES,
|
||||
help="Cooldown slices to use when estimating false accepts per hour.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--positive-skip-slices",
|
||||
type=int,
|
||||
default=DEFAULT_POSITIVE_SKIP_SLICES,
|
||||
help="Initial streaming slices to ignore when scoring positive examples.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cutoff-step",
|
||||
type=float,
|
||||
default=DEFAULT_CUTOFF_STEP,
|
||||
help="Cutoff increment to evaluate between cutoff-min and cutoff-max.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cutoff-min",
|
||||
type=float,
|
||||
default=DEFAULT_CUTOFF_MIN,
|
||||
help="Minimum cutoff to evaluate.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--cutoff-max",
|
||||
type=float,
|
||||
default=DEFAULT_CUTOFF_MAX,
|
||||
help="Maximum cutoff to evaluate.",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _parse_window_sizes(raw: str) -> list[int]:
|
||||
values = []
|
||||
for item in (raw or "").split(","):
|
||||
item = item.strip()
|
||||
if not item:
|
||||
continue
|
||||
value = int(item)
|
||||
if value < 1:
|
||||
raise ValueError("window sizes must be >= 1")
|
||||
values.append(value)
|
||||
if not values:
|
||||
raise ValueError("at least one window size is required")
|
||||
return sorted(set(values))
|
||||
|
||||
|
||||
def _moving_average(values: Sequence[float], window_size: int) -> np.ndarray:
|
||||
array = np.asarray(values, dtype=np.float32)
|
||||
if array.size == 0:
|
||||
return array
|
||||
if window_size <= 1:
|
||||
return array
|
||||
if array.size < window_size:
|
||||
return np.asarray([float(array.mean())], dtype=np.float32)
|
||||
cumsum = np.cumsum(np.insert(array, 0, 0.0))
|
||||
averaged = (cumsum[window_size:] - cumsum[:-window_size]) / float(window_size)
|
||||
return averaged.astype(np.float32)
|
||||
|
||||
|
||||
def _compute_false_accepts_per_hour(
|
||||
probabilities_per_track: Iterable[np.ndarray],
|
||||
cutoffs: np.ndarray,
|
||||
cooldown_slices: int,
|
||||
stride: int,
|
||||
step_seconds: float,
|
||||
) -> tuple[np.ndarray, float]:
|
||||
cutoffs = np.asarray(cutoffs, dtype=np.float32)
|
||||
false_accepts = np.zeros(cutoffs.shape[0], dtype=np.float64)
|
||||
duration_hours = 0.0
|
||||
|
||||
for track_probabilities in probabilities_per_track:
|
||||
if track_probabilities.size == 0:
|
||||
continue
|
||||
duration_hours += (
|
||||
len(track_probabilities) * stride * step_seconds / 3600.0
|
||||
)
|
||||
cooldown = np.full(cutoffs.shape[0], cooldown_slices, dtype=np.int32)
|
||||
for probability in track_probabilities:
|
||||
cooldown = np.maximum(cooldown - 1, 0)
|
||||
accepted = (cooldown == 0) & (probability > cutoffs)
|
||||
false_accepts += accepted.astype(np.float64)
|
||||
cooldown = np.where(accepted, cooldown_slices, cooldown)
|
||||
|
||||
if duration_hours <= 0:
|
||||
return np.full(cutoffs.shape[0], math.inf, dtype=np.float64), 0.0
|
||||
|
||||
return false_accepts / duration_hours, duration_hours
|
||||
|
||||
|
||||
def _select_best_candidate(
|
||||
candidates: list[dict[str, float]],
|
||||
target_faph: float,
|
||||
recall_margin: float = DEFAULT_RECALL_MARGIN,
|
||||
) -> tuple[dict[str, float], float]:
|
||||
if not candidates:
|
||||
raise ValueError("at least one calibration candidate is required")
|
||||
if recall_margin < 0:
|
||||
raise ValueError("recall margin must be >= 0")
|
||||
|
||||
fallback_limits = [
|
||||
target_faph,
|
||||
max(target_faph * 2.0, target_faph + 0.5),
|
||||
max(target_faph * 4.0, 2.0),
|
||||
]
|
||||
|
||||
def tier(candidate: dict[str, float]) -> int:
|
||||
for index, limit in enumerate(fallback_limits):
|
||||
if candidate["false_accepts_per_hour"] <= limit + 1e-9:
|
||||
return index
|
||||
return len(fallback_limits)
|
||||
|
||||
# Stay in the strictest false-accept tier that has a viable candidate. Within
|
||||
# that tier, keep candidates close to the best recall, then spend the allowed
|
||||
# recall margin on the lowest measured false-accept rate.
|
||||
best_tier = min(tier(candidate) for candidate in candidates)
|
||||
tier_candidates = [
|
||||
candidate for candidate in candidates if tier(candidate) == best_tier
|
||||
]
|
||||
best_recall = max(candidate["recall"] for candidate in tier_candidates)
|
||||
recall_floor = best_recall - recall_margin
|
||||
viable_candidates = [
|
||||
candidate
|
||||
for candidate in tier_candidates
|
||||
if candidate["recall"] >= recall_floor - 1e-12
|
||||
]
|
||||
|
||||
best = min(
|
||||
viable_candidates,
|
||||
key=lambda candidate: (
|
||||
candidate["false_accepts_per_hour"],
|
||||
-candidate["recall"],
|
||||
abs(candidate["sliding_window_size"] - PREFERRED_WINDOW_SIZE),
|
||||
-candidate["probability_cutoff"],
|
||||
),
|
||||
)
|
||||
|
||||
tier_index = tier(best)
|
||||
if tier_index < len(fallback_limits):
|
||||
return best, fallback_limits[tier_index]
|
||||
return best, float("inf")
|
||||
|
||||
|
||||
def _load_config(config_path: Path) -> dict:
|
||||
with config_path.open("r", encoding="utf-8") as handle:
|
||||
return yaml.load(handle.read(), Loader=yaml.Loader)
|
||||
|
||||
|
||||
def _load_eval_sets(
|
||||
handler: Any,
|
||||
config: dict,
|
||||
) -> tuple[str, str, list[np.ndarray], list[np.ndarray]]:
|
||||
for positive_mode, ambient_mode in (
|
||||
("validation", "validation_ambient"),
|
||||
("testing", "testing_ambient"),
|
||||
):
|
||||
positive_tracks, labels, _ = handler.get_data(
|
||||
positive_mode,
|
||||
batch_size=config["batch_size"],
|
||||
features_length=config["spectrogram_length"],
|
||||
truncation_strategy="none",
|
||||
)
|
||||
ambient_tracks, _, _ = handler.get_data(
|
||||
ambient_mode,
|
||||
batch_size=config["batch_size"],
|
||||
features_length=config["spectrogram_length"],
|
||||
truncation_strategy="none",
|
||||
)
|
||||
positives = [
|
||||
np.asarray(track)
|
||||
for track, label in zip(positive_tracks, labels)
|
||||
if bool(label)
|
||||
]
|
||||
ambient = [np.asarray(track) for track in ambient_tracks]
|
||||
if positives and ambient:
|
||||
return positive_mode, ambient_mode, positives, ambient
|
||||
raise RuntimeError(
|
||||
"No suitable validation/testing data was found for detector calibration."
|
||||
)
|
||||
|
||||
|
||||
def _predict_tracks(
|
||||
model: Any,
|
||||
tracks: Sequence[np.ndarray],
|
||||
label: str,
|
||||
) -> list[np.ndarray]:
|
||||
predictions: list[np.ndarray] = []
|
||||
total = len(tracks)
|
||||
print(f"→ Running streaming inference on {total} {label} track(s)")
|
||||
for index, track in enumerate(tracks, start=1):
|
||||
values = np.asarray(model.predict_spectrogram(track), dtype=np.float32)
|
||||
predictions.append(values)
|
||||
if index == total or index % 25 == 0:
|
||||
print(f" {label}: {index}/{total}")
|
||||
return predictions
|
||||
|
||||
|
||||
def main() -> int:
|
||||
from microwakeword.data import FeatureHandler
|
||||
from microwakeword.inference import Model
|
||||
|
||||
args = parse_args()
|
||||
window_sizes = _parse_window_sizes(args.window_sizes)
|
||||
if args.recall_margin < 0 or args.recall_margin > 1:
|
||||
raise ValueError("recall-margin must be between 0 and 1")
|
||||
if args.cutoff_step <= 0:
|
||||
raise ValueError("cutoff-step must be > 0")
|
||||
if args.cutoff_max < args.cutoff_min:
|
||||
raise ValueError("cutoff-max must be >= cutoff-min")
|
||||
|
||||
config_path = Path(args.training_config)
|
||||
model_path = Path(args.model)
|
||||
output_path = Path(args.output)
|
||||
|
||||
if not config_path.exists():
|
||||
raise FileNotFoundError(f"Training config not found: {config_path}")
|
||||
if not model_path.exists():
|
||||
raise FileNotFoundError(f"Streaming TFLite model not found: {model_path}")
|
||||
|
||||
cutoffs = np.arange(
|
||||
args.cutoff_min,
|
||||
args.cutoff_max + (args.cutoff_step / 2.0),
|
||||
args.cutoff_step,
|
||||
dtype=np.float32,
|
||||
)
|
||||
cutoffs = np.clip(cutoffs, 0.0, 1.0)
|
||||
cutoffs = np.unique(np.round(cutoffs, 4))
|
||||
|
||||
print("===== Detector Calibration =====")
|
||||
print(f"→ Model: {model_path}")
|
||||
print(f"→ Training config: {config_path}")
|
||||
print(
|
||||
f"→ Evaluating window sizes {window_sizes} with target <= "
|
||||
f"{args.target_faph:.2f} false accepts/hour"
|
||||
)
|
||||
print(
|
||||
f"→ Favoring lower false accepts within "
|
||||
f"{args.recall_margin:.2%} of the best recall"
|
||||
)
|
||||
|
||||
config = _load_config(config_path)
|
||||
config["flags"] = config.get("flags", {})
|
||||
handler = FeatureHandler(config)
|
||||
|
||||
positive_mode, ambient_mode, positive_tracks, ambient_tracks = _load_eval_sets(
|
||||
handler, config
|
||||
)
|
||||
|
||||
print(
|
||||
f"→ Using {positive_mode} positives ({len(positive_tracks)}) and "
|
||||
f"{ambient_mode} ambient tracks ({len(ambient_tracks)})"
|
||||
)
|
||||
|
||||
model = Model(str(model_path), stride=config["stride"])
|
||||
positive_predictions = _predict_tracks(model, positive_tracks, "positive")
|
||||
ambient_predictions = _predict_tracks(model, ambient_tracks, "ambient")
|
||||
|
||||
candidates: list[dict[str, float]] = []
|
||||
best_by_window: list[dict[str, float]] = []
|
||||
step_seconds = config["window_step_ms"] / 1000.0
|
||||
|
||||
for window_size in window_sizes:
|
||||
ambient_averages = [
|
||||
_moving_average(track, window_size) for track in ambient_predictions
|
||||
]
|
||||
positive_maxima = []
|
||||
for track in positive_predictions:
|
||||
search = (
|
||||
track[args.positive_skip_slices :]
|
||||
if track.size > args.positive_skip_slices
|
||||
else track
|
||||
)
|
||||
averaged = _moving_average(search, window_size)
|
||||
if averaged.size == 0:
|
||||
averaged = _moving_average(track, window_size)
|
||||
positive_maxima.append(float(np.max(averaged)) if averaged.size else 0.0)
|
||||
|
||||
positive_maxima_array = np.asarray(positive_maxima, dtype=np.float32)
|
||||
recall_by_cutoff = np.mean(
|
||||
positive_maxima_array[None, :] > cutoffs[:, None], axis=1
|
||||
)
|
||||
faph_by_cutoff, ambient_hours = _compute_false_accepts_per_hour(
|
||||
ambient_averages,
|
||||
cutoffs,
|
||||
args.cooldown_slices,
|
||||
stride=config["stride"],
|
||||
step_seconds=step_seconds,
|
||||
)
|
||||
|
||||
window_candidates = []
|
||||
for cutoff, recall, faph in zip(cutoffs, recall_by_cutoff, faph_by_cutoff):
|
||||
candidate = {
|
||||
"probability_cutoff": float(round(float(cutoff), 2)),
|
||||
"sliding_window_size": int(window_size),
|
||||
"recall": float(recall),
|
||||
"false_accepts_per_hour": float(faph),
|
||||
"ambient_hours": float(ambient_hours),
|
||||
}
|
||||
candidates.append(candidate)
|
||||
window_candidates.append(candidate)
|
||||
|
||||
best_window, _ = _select_best_candidate(
|
||||
window_candidates,
|
||||
args.target_faph,
|
||||
args.recall_margin,
|
||||
)
|
||||
best_by_window.append(best_window)
|
||||
print(
|
||||
" window={window}: cutoff={cutoff:.2f}; recall={recall:.2%}; "
|
||||
"ambient_faph={faph:.3f}".format(
|
||||
window=window_size,
|
||||
cutoff=best_window["probability_cutoff"],
|
||||
recall=best_window["recall"],
|
||||
faph=best_window["false_accepts_per_hour"],
|
||||
)
|
||||
)
|
||||
|
||||
best, selected_limit = _select_best_candidate(
|
||||
candidates,
|
||||
args.target_faph,
|
||||
args.recall_margin,
|
||||
)
|
||||
if best["false_accepts_per_hour"] > args.target_faph + 1e-9:
|
||||
print(
|
||||
"⚠️ No candidate met the target false accepts/hour budget; "
|
||||
"using the best fallback operating point."
|
||||
)
|
||||
|
||||
print(
|
||||
"✓ Selected cutoff={cutoff:.2f}, window={window}, recall={recall:.2%}, "
|
||||
"ambient_faph={faph:.3f}".format(
|
||||
cutoff=best["probability_cutoff"],
|
||||
window=best["sliding_window_size"],
|
||||
recall=best["recall"],
|
||||
faph=best["false_accepts_per_hour"],
|
||||
)
|
||||
)
|
||||
|
||||
output = {
|
||||
"probability_cutoff": best["probability_cutoff"],
|
||||
"sliding_window_size": best["sliding_window_size"],
|
||||
"target_false_accepts_per_hour": float(args.target_faph),
|
||||
"selected_false_accepts_per_hour_limit": (
|
||||
None if math.isinf(selected_limit) else float(selected_limit)
|
||||
),
|
||||
"selected_metrics": {
|
||||
"recall": round(best["recall"], 6),
|
||||
"false_accepts_per_hour": round(best["false_accepts_per_hour"], 6),
|
||||
"ambient_hours": round(best["ambient_hours"], 6),
|
||||
},
|
||||
"evaluation": {
|
||||
"positive_dataset": positive_mode,
|
||||
"ambient_dataset": ambient_mode,
|
||||
"positive_tracks": len(positive_tracks),
|
||||
"ambient_tracks": len(ambient_tracks),
|
||||
"cooldown_slices": int(args.cooldown_slices),
|
||||
"positive_skip_slices": int(args.positive_skip_slices),
|
||||
"window_sizes": window_sizes,
|
||||
"cutoff_min": round(float(cutoffs[0]), 4),
|
||||
"cutoff_max": round(float(cutoffs[-1]), 4),
|
||||
"cutoff_step": float(args.cutoff_step),
|
||||
"recall_margin": float(args.recall_margin),
|
||||
"preferred_window_size": PREFERRED_WINDOW_SIZE,
|
||||
},
|
||||
"per_window_best": best_by_window,
|
||||
"generated_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
output_path.write_text(json.dumps(output, indent=2) + "\n", encoding="utf-8")
|
||||
print(f"📝 Wrote calibration to {output_path}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -130,6 +130,73 @@ print(f" AudioSet complete ({ok} ok, {skipped} skipped, {len(audioset_bad)} fa
|
||||
EOF
|
||||
}
|
||||
|
||||
converter_from_dataset_api() {
|
||||
# shellcheck source=/dev/null
|
||||
source "${DATA_DIR}/.venv/bin/activate"
|
||||
|
||||
python - "${AUDIO16K_DIR}" <<-'EOF'
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import scipy.io.wavfile
|
||||
from datasets import load_dataset
|
||||
|
||||
def write_wav(dst: Path, data: np.ndarray, sr: int):
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
x = np.clip(data, -1.0, 1.0)
|
||||
scipy.io.wavfile.write(dst, sr, (x * 32767).astype(np.int16))
|
||||
|
||||
audioset_out = Path(sys.argv[1])
|
||||
|
||||
print(" AudioSet FLAC tarballs are unavailable; using Hugging Face datasets API instead.")
|
||||
dataset = load_dataset(
|
||||
"agkphysics/AudioSet",
|
||||
"balanced",
|
||||
split="train",
|
||||
streaming=True,
|
||||
)
|
||||
|
||||
audioset_bad = []
|
||||
ok = 0
|
||||
skipped = 0
|
||||
heartbeat_every = 250
|
||||
|
||||
for idx, sample in enumerate(dataset, start=1):
|
||||
try:
|
||||
video_id = str(sample.get("video_id") or f"audioset_{idx:06d}")
|
||||
outfile = audioset_out / f"{video_id}.wav"
|
||||
if outfile.exists():
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
audio = sample.get("audio") or {}
|
||||
y = np.asarray(audio.get("array"))
|
||||
sr = int(audio.get("sampling_rate") or 0)
|
||||
if y.size == 0 or sr <= 0:
|
||||
raise ValueError("missing decoded audio")
|
||||
if y.ndim > 1:
|
||||
y = np.mean(y, axis=-1)
|
||||
if sr != 16000:
|
||||
y = librosa.resample(y.astype(np.float32), orig_sr=sr, target_sr=16000)
|
||||
if y.size == 0:
|
||||
raise ValueError("empty audio")
|
||||
write_wav(outfile, y, 16000)
|
||||
ok += 1
|
||||
except Exception as exc:
|
||||
audioset_bad.append(f"{sample.get('video_id', idx)}:{exc}")
|
||||
|
||||
if idx == 1 or (idx % heartbeat_every) == 0:
|
||||
print(f" AudioSet API progress: {idx} clips processed (ok={ok}, skipped={skipped}, failed={len(audioset_bad)})")
|
||||
|
||||
if audioset_bad:
|
||||
(audioset_out / "audioset_corrupted_files.log").write_text("\n".join(audioset_bad))
|
||||
|
||||
print(f" AudioSet complete via datasets API ({ok} ok, {skipped} skipped, {len(audioset_bad)} failed)")
|
||||
EOF
|
||||
}
|
||||
|
||||
expected_filecount=$(get_total_filecount filecounts)
|
||||
actual_filecount=$(find "${AUDIO16K_DIR}" -name "*.wav" 2>/dev/null | wc -l) || :
|
||||
write_filecount=false
|
||||
@@ -139,40 +206,44 @@ if [ "${actual_filecount}" -ne 0 ] ; then
|
||||
echo " Existing ${AUDIO16K_DIR} present (${actual_filecount} wav); skipping extract/convert"
|
||||
else
|
||||
dl=$(find_rev)
|
||||
[ -n "$dl" ] || { echo " Could not locate an AudioSet revision with FLAC tarballs still present on HF." ; exit 1 ; }
|
||||
rev=${dl%%,*}
|
||||
pattern=${dl##*,}
|
||||
if [ -z "$dl" ] ; then
|
||||
rm -rf "${AUDIO16K_DIR}/audioset_corrupted_files.log" || :
|
||||
converter_from_dataset_api
|
||||
else
|
||||
rev=${dl%%,*}
|
||||
pattern=${dl##*,}
|
||||
|
||||
echo " Checking 10 tarballs"
|
||||
for i in {0..9} ; do
|
||||
fname="downloads/bal_train0${i}.tar"
|
||||
if [ ! -f "${fname}" ] ; then
|
||||
echo " Downloading bal_train0${i}.tar"
|
||||
url="${AUDIO_URL}/${rev}/${pattern}${i}.tar"
|
||||
curl -L -s --fail "${url}" -o "${fname}" || { echo "Could not fetch ${fname} at rev ${rev}; continuing." ; continue ; }
|
||||
echo " Checking 10 tarballs"
|
||||
for i in {0..9} ; do
|
||||
fname="downloads/bal_train0${i}.tar"
|
||||
if [ ! -f "${fname}" ] ; then
|
||||
echo " Downloading bal_train0${i}.tar"
|
||||
url="${AUDIO_URL}/${rev}/${pattern}${i}.tar"
|
||||
curl -L -s --fail "${url}" -o "${fname}" || { echo "Could not fetch ${fname} at rev ${rev}; continuing." ; continue ; }
|
||||
fi
|
||||
|
||||
tarball_filecount=$(tar -tvf "${fname}" | wc -l )
|
||||
filecounts["bal_train0${i}.tar"]=${tarball_filecount}
|
||||
write_filecount=true
|
||||
|
||||
echo " Untarring bal_train0${i}.tar"
|
||||
tar -xf "${fname}" -C "${AUDIO_DIR}"
|
||||
if "${CLEANUP_ARCHIVES}" && [ -f "${fname}" ] ; then
|
||||
echo " Cleaning up bal_train0${i}.tar"
|
||||
rm -rf "${fname}"
|
||||
fi
|
||||
done
|
||||
|
||||
rm -rf "${AUDIO16K_DIR}/audioset_corrupted_files.log" || :
|
||||
converter
|
||||
|
||||
# Recompute counts and warn (but do not fail)
|
||||
expected_filecount=$(get_total_filecount filecounts)
|
||||
actual_filecount=$(find "${AUDIO16K_DIR}" -name "*.wav" 2>/dev/null | wc -l) || :
|
||||
if [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
echo " Converted file count(${actual_filecount}) != expected file count(${expected_filecount})" >&2
|
||||
echo " WARNING: mismatch is expected if some AudioSet files are corrupted; continuing." >&2
|
||||
fi
|
||||
|
||||
tarball_filecount=$(tar -tvf "${fname}" | wc -l )
|
||||
filecounts["bal_train0${i}.tar"]=${tarball_filecount}
|
||||
write_filecount=true
|
||||
|
||||
echo " Untarring bal_train0${i}.tar"
|
||||
tar -xf "${fname}" -C "${AUDIO_DIR}"
|
||||
if "${CLEANUP_ARCHIVES}" && [ -f "${fname}" ] ; then
|
||||
echo " Cleaning up bal_train0${i}.tar"
|
||||
rm -rf "${fname}"
|
||||
fi
|
||||
done
|
||||
|
||||
rm -rf "${AUDIO16K_DIR}/audioset_corrupted_files.log" || :
|
||||
converter
|
||||
|
||||
# Recompute counts and warn (but do not fail)
|
||||
expected_filecount=$(get_total_filecount filecounts)
|
||||
actual_filecount=$(find "${AUDIO16K_DIR}" -name "*.wav" 2>/dev/null | wc -l) || :
|
||||
if [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
echo " Converted file count(${actual_filecount}) != expected file count(${expected_filecount})" >&2
|
||||
echo " WARNING: mismatch is expected if some AudioSet files are corrupted; continuing." >&2
|
||||
fi
|
||||
fi
|
||||
|
||||
@@ -196,4 +267,4 @@ if "${CLEANUP_INTERMEDIATE_FILES}" && [ -d "${AUDIO_DIR}" ] ; then
|
||||
fi
|
||||
|
||||
echo " Audioset complete"
|
||||
exit 0
|
||||
exit 0
|
||||
|
||||
112
cli/setup_blackwell_venv
Executable file
112
cli/setup_blackwell_venv
Executable file
@@ -0,0 +1,112 @@
|
||||
#!/bin/bash
|
||||
set -euo pipefail
|
||||
|
||||
PROGDIR="$(dirname "$(realpath "$0")")"
|
||||
ROOTDIR="$(dirname "${PROGDIR}")"
|
||||
|
||||
KNOWN_ARGS=( data-dir force python )
|
||||
source "${PROGDIR}/shell.functions"
|
||||
|
||||
if [ ${#UNKNOWN_ARGS[@]} -gt 0 ] ; then
|
||||
echo "Unknown argument(s): ${UNKNOWN_ARGS[*]}" >&2
|
||||
HELP=true
|
||||
fi
|
||||
|
||||
if [ "${HELP}" == "true" ] ; then
|
||||
cat <<EOF >&2
|
||||
Usage: setup_blackwell_venv [ --data-dir=/data ] [ --force ] [ --python=python3.13 ]
|
||||
|
||||
Creates /data/.venv-blackwell for RTX 50 / Blackwell TensorFlow training.
|
||||
Sample generation and augmentation continue to use /data/.venv.
|
||||
|
||||
Environment overrides:
|
||||
MWW_BLACKWELL_TF_WHEEL_URL: TensorFlow Blackwell wheel URL.
|
||||
|
||||
EOF
|
||||
exit 1
|
||||
fi
|
||||
|
||||
[ -n "${DATA_DIR}" ] && DATA_DIR="$(realpath "${DATA_DIR}")"
|
||||
[ -d "${DATA_DIR}" ] || {
|
||||
echo "Data directory '${DATA_DIR}' doesn't exist." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
PYTHON="${PYTHON:-python3.13}"
|
||||
VENV="${DATA_DIR}/.venv-blackwell"
|
||||
MARKER="${VENV}/.mww-blackwell-venv"
|
||||
TF_WHEEL_URL="${MWW_BLACKWELL_TF_WHEEL_URL:-https://github.com/chivitiH/tensorflow-blackwell-python313/releases/download/v2.22.0-selfbuilt/tensorflow-2.22.0.dev0+selfbuilt-cp313-cp313-linux_x86_64.whl}"
|
||||
|
||||
if ! command -v "${PYTHON}" >/dev/null 2>&1 ; then
|
||||
echo "Python 3.13 is required for the Blackwell TensorFlow wheel. Missing: ${PYTHON}" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if [ "${FORCE:-false}" != "true" ] && [ -x "${VENV}/bin/python" ] && [ -f "${MARKER}" ] ; then
|
||||
echo " Blackwell TensorFlow venv found (skipping setup_blackwell_venv)"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
echo "===== Setting up Blackwell TensorFlow environment ${VENV} ====="
|
||||
rm -rf "${VENV}" || :
|
||||
"${PYTHON}" -m venv --upgrade-deps "${VENV}"
|
||||
source "${VENV}/bin/activate"
|
||||
|
||||
export PIP_PROGRESS_BAR=off
|
||||
export PIP_NO_COLOR=1
|
||||
export PIP_QUIET=0
|
||||
|
||||
pip_install() {
|
||||
if $VERBOSE ; then
|
||||
pip install "$@" || return 1
|
||||
else
|
||||
{ pip install "$@" || return 1 ; } | stdbuf -i0 -o0 tr -d '[:print:]' | stdbuf -i0 -o0 tr '\n' '.'
|
||||
fi
|
||||
echo
|
||||
}
|
||||
|
||||
echo " ===== Installing Blackwell TensorFlow wheel ====="
|
||||
pip_install --upgrade pip setuptools wheel
|
||||
pip_install "${TF_WHEEL_URL}"
|
||||
|
||||
echo " ===== Installing microWakeWord training dependencies ====="
|
||||
pip_install \
|
||||
audiomentations \
|
||||
audio_metadata \
|
||||
datasets \
|
||||
mmap_ninja \
|
||||
pymicro-features \
|
||||
pyyaml \
|
||||
webrtcvad-wheels \
|
||||
ai-edge-litert \
|
||||
numpy-minmax \
|
||||
numpy-rms \
|
||||
absl-py \
|
||||
"numpy==2.3.5"
|
||||
|
||||
echo " ===== Checking microwakeword ====="
|
||||
MWW="${DATA_DIR}/tools/microWakeWord"
|
||||
if [ ! -d "${MWW}" ] || [ -n "$(git -C "${MWW}" status --porcelain 2>/dev/null || true)" ] ; then
|
||||
rm -rf "${MWW}" || :
|
||||
mkdir -p "${DATA_DIR}/tools"
|
||||
echo " Cloning micro-wake-word to ${DATA_DIR}/tools"
|
||||
git clone https://github.com/TaterTotterson/micro-wake-word "${MWW}" &>/dev/null
|
||||
fi
|
||||
echo " Installing microwakeword into Blackwell venv"
|
||||
pip_install --no-deps -e "${MWW}"
|
||||
|
||||
echo " ===== Testing Blackwell TensorFlow environment ====="
|
||||
"${VENV}/bin/python" - <<'PY'
|
||||
import tensorflow as tf
|
||||
from ai_edge_litert.interpreter import Interpreter
|
||||
from microwakeword.data import FeatureHandler
|
||||
from microwakeword.inference import Model
|
||||
|
||||
print("TensorFlow:", tf.__version__)
|
||||
print("CUDA build:", tf.test.is_built_with_cuda())
|
||||
print("GPU:", tf.config.list_physical_devices("GPU"))
|
||||
print("microWakeWord Blackwell imports available")
|
||||
PY
|
||||
|
||||
touch "${MARKER}"
|
||||
echo "Blackwell TensorFlow environment ready: ${VENV}"
|
||||
@@ -27,8 +27,11 @@ cd "${DATA_DIR}/training_datasets"
|
||||
|
||||
echo "***** Checking FMA *****"
|
||||
|
||||
AUDIO_URL="https://huggingface.co/datasets/mchl914/fma_xsmall/resolve/main/fma_xs.zip"
|
||||
AUDIO_ZIPFILE="fma_xs.zip"
|
||||
AUDIO_URLS=(
|
||||
"https://os.unil.cloud.switch.ch/fma/fma_small.zip"
|
||||
"https://huggingface.co/datasets/mchl914/fma_xsmall/resolve/main/fma_xs.zip"
|
||||
)
|
||||
AUDIO_ZIPFILE="fma_small.zip"
|
||||
AUDIO_ZIP="./downloads/${AUDIO_ZIPFILE}"
|
||||
AUDIO_DIR="fma"
|
||||
mkdir -p "${AUDIO_DIR}" || :
|
||||
@@ -81,6 +84,52 @@ EOF
|
||||
|
||||
}
|
||||
|
||||
extract_zip_with_python() {
|
||||
local zip_path="$1"
|
||||
local dest_dir="$2"
|
||||
|
||||
"${DATA_DIR}/.venv/bin/python" - "${zip_path}" "${dest_dir}" <<-'EOF'
|
||||
import sys
|
||||
import zipfile
|
||||
from pathlib import Path
|
||||
from tqdm import tqdm
|
||||
|
||||
zip_path = Path(sys.argv[1])
|
||||
dest_dir = Path(sys.argv[2])
|
||||
|
||||
if (not zip_path.exists()) or zip_path.stat().st_size == 0:
|
||||
raise SystemExit(f"Archive missing or empty: {zip_path}")
|
||||
|
||||
with zipfile.ZipFile(zip_path, "r") as zf:
|
||||
members = zf.infolist()
|
||||
size_gb = zip_path.stat().st_size / (1024 ** 3)
|
||||
print(f" Extracting {zip_path.name} ({len(members)} entries, {size_gb:.1f} GiB)...")
|
||||
for member in tqdm(members, desc=" FMA zip extract", unit="file"):
|
||||
zf.extract(member, dest_dir)
|
||||
EOF
|
||||
}
|
||||
|
||||
download_with_fallbacks() {
|
||||
local output="$1"
|
||||
shift
|
||||
local urls=( "$@" )
|
||||
local rc=1
|
||||
|
||||
for url in "${urls[@]}" ; do
|
||||
for attempt in 1 2 3 4 ; do
|
||||
curl -sfL "${url}" -o "${output}" && [ -s "${output}" ] && return 0
|
||||
rc=$?
|
||||
rm -f "${output}" || :
|
||||
if [ "${attempt}" -lt 4 ] ; then
|
||||
echo " Retry ${attempt}/3 after download failure"
|
||||
sleep $(( attempt * 2 ))
|
||||
fi
|
||||
done
|
||||
done
|
||||
|
||||
return "${rc}"
|
||||
}
|
||||
|
||||
expected_filecount=${filecounts[${AUDIO_ZIPFILE}]}
|
||||
actual_filecount=$(find ${AUDIO16K_DIR} -name '*.wav' 2>/dev/null | wc -l) || :
|
||||
write_filecount=false
|
||||
@@ -92,13 +141,16 @@ else
|
||||
if [ "${actual_filecount}" -eq 0 ] || [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
if [ ! -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Downloading ${AUDIO_ZIPFILE}"
|
||||
curl -sfL "${AUDIO_URL}" -o "${AUDIO_ZIP}"
|
||||
download_with_fallbacks "${AUDIO_ZIP}" "${AUDIO_URLS[@]}" || {
|
||||
echo " Failed to download ${AUDIO_ZIPFILE} from all configured sources." >&2
|
||||
exit 1
|
||||
}
|
||||
fi
|
||||
|
||||
rm -rf "${AUDIO_DIR}" || :
|
||||
mkdir "${AUDIO_DIR}"
|
||||
echo " Unzipping ${AUDIO_ZIPFILE}"
|
||||
unzip -q -d "${AUDIO_DIR}" "${AUDIO_ZIP}"
|
||||
echo " Extracting ${AUDIO_ZIPFILE}"
|
||||
extract_zip_with_python "${AUDIO_ZIP}" "${AUDIO_DIR}"
|
||||
fi
|
||||
if "${CLEANUP_ARCHIVES}" && [ -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Cleaning up ${AUDIO_ZIPFILE}"
|
||||
@@ -128,4 +180,3 @@ fi
|
||||
|
||||
echo " FMA complete"
|
||||
exit 0
|
||||
|
||||
|
||||
@@ -25,9 +25,9 @@ fi
|
||||
mkdir -p "${DATA_DIR}/training_datasets/downloads" || :
|
||||
cd "${DATA_DIR}/training_datasets"
|
||||
|
||||
AUDIO_URL="https://mcdermottlab.mit.edu/Reverb/IRMAudio/Audio.zip"
|
||||
AUDIO_ZIPFILE="MIT_RIR_Audio.zip"
|
||||
AUDIO_ZIP="./downloads/${AUDIO_ZIPFILE}"
|
||||
HF_RIR_REPO_ID="TaterTotterson/MIT_environmental_impulse_responses"
|
||||
HF_RIR_API_URL="https://huggingface.co/api/datasets/${HF_RIR_REPO_ID}"
|
||||
HF_RIR_SOURCE_KEY="hf_mit_environmental_impulse_responses"
|
||||
AUDIO_DIR="./mit_rirs"
|
||||
mkdir -p "${AUDIO_DIR}" || :
|
||||
AUDIO16K_DIR="./mit_rirs_16k"
|
||||
@@ -35,10 +35,92 @@ mkdir -p "${AUDIO16K_DIR}" || :
|
||||
AUDIO_FILECOUNT="./downloads/mit_rir_filecount"
|
||||
AUDIO_IN_GLOB="*.wav"
|
||||
|
||||
declare -A filecounts=( [${AUDIO_ZIPFILE}]=0 )
|
||||
declare -A filecounts=( [${HF_RIR_SOURCE_KEY}]=0 )
|
||||
get_filecounts filecounts "${AUDIO_FILECOUNT}"
|
||||
|
||||
echo "===== Checking MIT_RIR ====="
|
||||
echo "===== Checking MIT environmental RIRs ====="
|
||||
|
||||
download_hf_mit_rirs() {
|
||||
source ${DATA_DIR}/.venv/bin/activate
|
||||
python - "${HF_RIR_REPO_ID}" "${HF_RIR_API_URL}" "${AUDIO_DIR}" <<-'EOF'
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
repo_id = sys.argv[1]
|
||||
api_url = sys.argv[2]
|
||||
audio_dir = Path(sys.argv[3])
|
||||
audio_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
request = urllib.request.Request(api_url, headers={"User-Agent": "WakeWordTrainer/1.0"})
|
||||
with urllib.request.urlopen(request, timeout=30) as response:
|
||||
metadata = json.loads(response.read().decode("utf-8"))
|
||||
|
||||
files = sorted(
|
||||
sibling.get("rfilename", "")
|
||||
for sibling in metadata.get("siblings", [])
|
||||
if str(sibling.get("rfilename", "")).startswith("16khz/")
|
||||
and str(sibling.get("rfilename", "")).lower().endswith(".wav")
|
||||
)
|
||||
if not files:
|
||||
raise SystemExit("Hugging Face MIT RIR dataset did not list any 16khz WAV files")
|
||||
|
||||
print(f" Found {len(files)} MIT environmental RIR files on Hugging Face mirror", flush=True)
|
||||
downloaded = 0
|
||||
skipped = 0
|
||||
|
||||
def download_file(url: str, target: Path, rel: str):
|
||||
tmp = target.with_suffix(target.suffix + ".incomplete")
|
||||
for attempt in range(1, 4):
|
||||
try:
|
||||
if tmp.exists():
|
||||
tmp.unlink()
|
||||
with urllib.request.urlopen(url, timeout=30) as response:
|
||||
with tmp.open("wb") as out:
|
||||
while True:
|
||||
chunk = response.read(1024 * 64)
|
||||
if not chunk:
|
||||
break
|
||||
out.write(chunk)
|
||||
if not tmp.exists() or tmp.stat().st_size == 0:
|
||||
raise RuntimeError("empty download")
|
||||
tmp.replace(target)
|
||||
return
|
||||
except Exception as exc:
|
||||
if tmp.exists():
|
||||
tmp.unlink()
|
||||
if attempt == 3:
|
||||
raise RuntimeError(f"download failed for {rel}: {exc}") from exc
|
||||
print(f" Retry {attempt}/2 for {rel}: {exc}", flush=True)
|
||||
time.sleep(2 * attempt)
|
||||
|
||||
total = len(files)
|
||||
for idx, rel in enumerate(files, start=1):
|
||||
target = audio_dir / rel
|
||||
if target.exists() and target.stat().st_size > 0:
|
||||
skipped += 1
|
||||
if idx == 1 or idx % 25 == 0 or idx == total:
|
||||
print(f" MIT RIR download progress: {idx}/{total} files ({downloaded} downloaded, {skipped} reused)", flush=True)
|
||||
continue
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
encoded = urllib.parse.quote(rel, safe="/")
|
||||
url = f"https://huggingface.co/datasets/{repo_id}/resolve/main/{encoded}"
|
||||
if idx == 1 or idx % 25 == 0 or idx == total:
|
||||
print(f" Downloading MIT RIR {idx}/{total}: {rel}", flush=True)
|
||||
download_file(url, target, rel)
|
||||
if not target.exists() or target.stat().st_size == 0:
|
||||
raise SystemExit(f"download failed for {rel}")
|
||||
downloaded += 1
|
||||
if idx == 1 or idx % 25 == 0 or idx == total:
|
||||
print(f" MIT RIR download progress: {idx}/{total} files ({downloaded} downloaded, {skipped} reused)", flush=True)
|
||||
|
||||
print(f" Hugging Face MIT environmental RIR download complete ({downloaded} downloaded, {skipped} reused)", flush=True)
|
||||
print(f" MIT environmental RIR files available: {len(files)}", flush=True)
|
||||
EOF
|
||||
}
|
||||
|
||||
converter() {
|
||||
source ${DATA_DIR}/.venv/bin/activate
|
||||
@@ -58,9 +140,9 @@ rir_out = Path(sys.argv[2])
|
||||
|
||||
waves = list(rir_in.rglob("*.wav"))
|
||||
try:
|
||||
print(" MIT RIR normalizing to 16k…")
|
||||
print(" MIT environmental RIR normalizing to 16k…")
|
||||
# Normalize to 16k mono
|
||||
for p in tqdm(waves, desc=" MIT_RIR (resample 16k mono)"):
|
||||
for p in tqdm(waves, desc=" MIT environmental RIR (resample 16k mono)"):
|
||||
outfile = Path(rir_out / p.name)
|
||||
if outfile.exists():
|
||||
continue
|
||||
@@ -70,14 +152,14 @@ try:
|
||||
if sr != 16000:
|
||||
a, _ = librosa.load(p, sr=16000, mono=True)
|
||||
write_wav(outfile, a, 16000)
|
||||
print(" MIT RIR normalization complete")
|
||||
print(" MIT environmental RIR normalization complete")
|
||||
except Exception as e2:
|
||||
print(f" MIT RIR fallback failed: {e2}")
|
||||
print(f" MIT environmental RIR preparation failed: {e2}")
|
||||
raise
|
||||
EOF
|
||||
}
|
||||
|
||||
expected_filecount=${filecounts[${AUDIO_ZIPFILE}]}
|
||||
expected_filecount=${filecounts[${HF_RIR_SOURCE_KEY}]}
|
||||
actual_filecount=$(find "${AUDIO16K_DIR}" -name '*.wav' 2>/dev/null | wc -l) || :
|
||||
write_filecount=false
|
||||
|
||||
@@ -85,24 +167,16 @@ if [ "${actual_filecount}" -ne 0 ] && [ "${actual_filecount}" -eq "${expected_fi
|
||||
echo " Existing ${AUDIO16K_DIR} valid"
|
||||
else
|
||||
actual_filecount=$(find "${AUDIO_DIR}" -name "${AUDIO_IN_GLOB}" 2>/dev/null | wc -l) || :
|
||||
if [ "${actual_filecount}" -eq 0 ] || [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
if [ ! -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Downloading ${AUDIO_ZIPFILE}"
|
||||
curl -sfL "${AUDIO_URL}" -o "${AUDIO_ZIP}"
|
||||
fi
|
||||
|
||||
if [ "${actual_filecount}" -eq 0 ] || [ "${expected_filecount}" -eq 0 ] || [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
rm -rf "${AUDIO_DIR}" || :
|
||||
echo " Unzipping ${AUDIO_ZIPFILE}"
|
||||
unzip -u -q -d "${AUDIO_DIR}" "${AUDIO_ZIP}"
|
||||
fi
|
||||
if "${CLEANUP_ARCHIVES}" && [ -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Cleaning up ${AUDIO_ZIPFILE}"
|
||||
rm -rf "${AUDIO_ZIP}"
|
||||
mkdir -p "${AUDIO_DIR}" || :
|
||||
echo " Downloading MIT environmental impulse responses from Hugging Face mirror"
|
||||
download_hf_mit_rirs
|
||||
fi
|
||||
|
||||
converter
|
||||
actual_filecount=$(find "${AUDIO16K_DIR}" -name "*.wav" 2>/dev/null | wc -l) || :
|
||||
filecounts[${AUDIO_ZIPFILE}]="${actual_filecount}"
|
||||
filecounts[${HF_RIR_SOURCE_KEY}]="${actual_filecount}"
|
||||
write_filecount=true
|
||||
fi
|
||||
|
||||
@@ -110,15 +184,10 @@ if ${write_filecount} ; then
|
||||
write_filecounts filecounts "${AUDIO_FILECOUNT}"
|
||||
fi
|
||||
|
||||
if "${CLEANUP_ARCHIVES}" && [ -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Cleaning up ${AUDIO_ZIPFILE}"
|
||||
rm -rf "${AUDIO_ZIP}"
|
||||
fi
|
||||
|
||||
if "${CLEANUP_INTERMEDIATE_FILES}" && [ -d "${AUDIO_DIR}" ]; then
|
||||
echo " Cleaning up ${AUDIO_DIR}"
|
||||
rm -rf "${AUDIO_DIR}"
|
||||
fi
|
||||
|
||||
echo " MIT_RIR complete"
|
||||
echo " MIT environmental RIRs complete"
|
||||
exit 0
|
||||
|
||||
@@ -242,29 +242,7 @@ if [ ! -s "${MODEL_FILE}.json" ] ; then
|
||||
curl -sfL "${MODEL_URL}.json" -o "${MODEL_FILE}.json"
|
||||
fi
|
||||
|
||||
# --- Dutch ONNX voices (single-speaker, used with --language=nl) ---
|
||||
# Working Dutch voices: pim, ronnie (nl_NL) and nathalie (nl_BE).
|
||||
# nl_NL-mls-medium is intentionally excluded (known Piper issue: outputs gibberish).
|
||||
HF_VOICES="https://huggingface.co/rhasspy/piper-voices/resolve/main"
|
||||
declare -a NL_VOICES=(
|
||||
"nl/nl_NL/pim/medium/nl_NL-pim-medium"
|
||||
"nl/nl_NL/ronnie/medium/nl_NL-ronnie-medium"
|
||||
"nl/nl_BE/nathalie/medium/nl_BE-nathalie-medium"
|
||||
)
|
||||
echo " ===== Checking Dutch Piper voices ====="
|
||||
for voice_path in "${NL_VOICES[@]}" ; do
|
||||
voice_name="$(basename "${voice_path}")"
|
||||
onnx_file="${VOICES_DIR}/${voice_name}.onnx"
|
||||
json_file="${VOICES_DIR}/${voice_name}.onnx.json"
|
||||
if [ ! -f "${onnx_file}" ] ; then
|
||||
echo " Downloading ${voice_name}.onnx"
|
||||
curl -sfL "${HF_VOICES}/${voice_path}.onnx?download=true" -o "${onnx_file}"
|
||||
fi
|
||||
if [ ! -f "${json_file}" ] ; then
|
||||
echo " Downloading ${voice_name}.onnx.json"
|
||||
curl -sfL "${HF_VOICES}/${voice_path}.onnx.json?download=true" -o "${json_file}"
|
||||
fi
|
||||
done
|
||||
echo " Non-English Piper voices will be downloaded on demand for the selected language."
|
||||
|
||||
${GPU} && onnxgpu='-gpu[cuda]' || onnxgpu=""
|
||||
echo " ===== Installing onnxruntime${onnxgpu} ====="
|
||||
|
||||
@@ -103,7 +103,8 @@ else
|
||||
if [ "${actual_filecount}" -eq 0 ] || [ "${actual_filecount}" -ne "${expected_filecount}" ] ; then
|
||||
if [ ! -f "${AUDIO_ZIP}" ] ; then
|
||||
echo " Downloading ${AUDIO_ZIPFILE}"
|
||||
curl -sfL "${AUDIO_URL}" -o "${AUDIO_ZIP}"
|
||||
curl -fL --progress-bar "${AUDIO_URL}" -o "${AUDIO_ZIP}" \
|
||||
2> >(tr '\r' '\n' >&2)
|
||||
fi
|
||||
|
||||
rm -rf "${AUDIO_DIR}" || :
|
||||
|
||||
@@ -18,6 +18,8 @@ parser.add_argument("--output-dir", type=str, help="Wake word output dir. Defaul
|
||||
# Personal inputs/outputs (NEW)
|
||||
parser.add_argument("--personal-dir", type=str, help="Personal WAV dir. Default: <data-dir>/personal_samples", required=False)
|
||||
parser.add_argument("--personal-output-dir", type=str, help="Personal features output dir. Default: <data-dir>/work/personal_augmented_features", required=False)
|
||||
parser.add_argument("--negative-dir", type=str, help="Reviewed negative WAV dir. Default: <data-dir>/negative_samples", required=False)
|
||||
parser.add_argument("--negative-output-dir", type=str, help="Reviewed negative features output dir. Default: <data-dir>/work/reviewed_negative_features", required=False)
|
||||
|
||||
# Dataset dirs
|
||||
parser.add_argument("--mit-rirs-16k-dir", type=str, help="MIT RIR input directory. Default: <data-dir>/training_datasets/mit_rirs_16k", required=False)
|
||||
@@ -57,6 +59,17 @@ if not args.personal_output_dir:
|
||||
else:
|
||||
args.personal_output_dir = os.path.realpath(args.personal_output_dir)
|
||||
|
||||
# Reviewed negative defaults
|
||||
if not args.negative_dir:
|
||||
args.negative_dir = os.path.join(args.data_dir, "negative_samples")
|
||||
else:
|
||||
args.negative_dir = os.path.realpath(args.negative_dir)
|
||||
|
||||
if not args.negative_output_dir:
|
||||
args.negative_output_dir = os.path.join(work_dir, "reviewed_negative_features")
|
||||
else:
|
||||
args.negative_output_dir = os.path.realpath(args.negative_output_dir)
|
||||
|
||||
# Dataset defaults
|
||||
if not args.mit_rirs_16k_dir:
|
||||
args.mit_rirs_16k_dir = os.path.join(args.data_dir, "training_datasets", "mit_rirs_16k")
|
||||
@@ -205,7 +218,7 @@ def bind_wav_generator(clips_obj: Clips, wav_dir: str):
|
||||
|
||||
clips_obj.audio_generator = types.MethodType(audio_generator_from_wavs, clips_obj)
|
||||
|
||||
def generate_feature_set(input_wav_dir: str, out_root_dir: str, label: str):
|
||||
def generate_feature_set(input_wav_dir: str, out_root_dir: str, label: str, *, remove_silence: bool = True):
|
||||
files = glob.glob(os.path.join(input_wav_dir, "*.wav"))
|
||||
if not files:
|
||||
print(f"ℹ️ No WAVs found for {label} in: {input_wav_dir} (skipping)")
|
||||
@@ -218,7 +231,7 @@ def generate_feature_set(input_wav_dir: str, out_root_dir: str, label: str):
|
||||
input_directory=input_wav_dir,
|
||||
file_pattern="*.wav",
|
||||
max_clip_duration_s=5,
|
||||
remove_silence=True,
|
||||
remove_silence=remove_silence,
|
||||
random_split_seed=10,
|
||||
split_count=0.1,
|
||||
)
|
||||
@@ -263,9 +276,12 @@ def generate_feature_set(input_wav_dir: str, out_root_dir: str, label: str):
|
||||
# Wake word generated/TTS features (existing behavior)
|
||||
generate_feature_set(args.input_dir, args.output_dir, "generated")
|
||||
|
||||
# Personal features (NEW)
|
||||
# Personal features
|
||||
generate_feature_set(args.personal_dir, args.personal_output_dir, "personal")
|
||||
|
||||
# Reviewed false-positive / hard-negative features
|
||||
generate_feature_set(args.negative_dir, args.negative_output_dir, "reviewed negatives", remove_silence=False)
|
||||
|
||||
END_TIME = datetime.now(timezone.utc).replace(microsecond=0)
|
||||
et = END_TIME - START_TIME
|
||||
print(f"\n{'=' * 80}")
|
||||
|
||||
@@ -84,6 +84,21 @@ if [ "${IS_BLACKWELL}" = "true" ]; then
|
||||
echo "ℹ️ Using GPU compatibility retries; CPU fallback is ${ALLOW_CPU_FALLBACK} (override with MWW_ALLOW_CPU_FALLBACK=true|false)."
|
||||
fi
|
||||
|
||||
BLACKWELL_TF_MODE="${MWW_BLACKWELL_TF:-auto}"
|
||||
BLACKWELL_TF_REQUIRED="false"
|
||||
BLACKWELL_TF_ACTIVE="false"
|
||||
case "${BLACKWELL_TF_MODE,,}" in
|
||||
1|true|yes|on|required)
|
||||
BLACKWELL_TF_REQUIRED="true"
|
||||
;;
|
||||
0|false|no|off|disabled)
|
||||
BLACKWELL_TF_MODE="disabled"
|
||||
;;
|
||||
*)
|
||||
BLACKWELL_TF_MODE="auto"
|
||||
;;
|
||||
esac
|
||||
|
||||
# Enable driver-side PTX JIT fallback when ptxas/nvlink are unavailable.
|
||||
if [ -z "${XLA_FLAGS:-}" ]; then
|
||||
export XLA_FLAGS="--xla_gpu_unsafe_fallback_to_driver_on_ptxas_not_found"
|
||||
@@ -111,6 +126,16 @@ else
|
||||
echo "ℹ️ No personal features found at ${PERSONAL_FEATURES_DIR}/training (continuing without personal weighting)"
|
||||
fi
|
||||
|
||||
# Reviewed false-positive features are optional hard negatives.
|
||||
REVIEWED_NEGATIVE_FEATURES_DIR="${WORK_DIR}/reviewed_negative_features"
|
||||
HAS_REVIEWED_NEGATIVE="false"
|
||||
if [ -d "${REVIEWED_NEGATIVE_FEATURES_DIR}/training" ] ; then
|
||||
HAS_REVIEWED_NEGATIVE="true"
|
||||
echo "✅ Found reviewed negative features: ${REVIEWED_NEGATIVE_FEATURES_DIR}/training (will weight as hard negatives)"
|
||||
else
|
||||
echo "ℹ️ No reviewed negative features found at ${REVIEWED_NEGATIVE_FEATURES_DIR}/training (continuing with stock negatives)"
|
||||
fi
|
||||
|
||||
cd "${WORK_DIR}"
|
||||
|
||||
echo "===== Starting ${TRAINING_STEPS} training steps ====="
|
||||
@@ -133,6 +158,7 @@ features:
|
||||
truth: true
|
||||
type: mmap
|
||||
__PERSONAL_FEATURE_MARKER__
|
||||
__REVIEWED_NEGATIVE_FEATURE_MARKER__
|
||||
- features_dir: __NEG_SPEECH__
|
||||
penalty_weight: 1.0
|
||||
sampling_weight: 12.0
|
||||
@@ -208,9 +234,51 @@ else
|
||||
sed -i -e "/__PERSONAL_FEATURE_MARKER__/d" "${YAML_PATH}"
|
||||
fi
|
||||
|
||||
# Insert/remove reviewed hard-negative block
|
||||
if [ "${HAS_REVIEWED_NEGATIVE}" = "true" ]; then
|
||||
reviewed_negative_block="$(cat <<EOF
|
||||
- features_dir: ${REVIEWED_NEGATIVE_FEATURES_DIR}
|
||||
penalty_weight: 1.25
|
||||
sampling_weight: 8.0
|
||||
truncation_strategy: random
|
||||
truth: false
|
||||
type: mmap
|
||||
EOF
|
||||
)"
|
||||
perl -0777 -i -pe "s#__REVIEWED_NEGATIVE_FEATURE_MARKER__#${reviewed_negative_block}#g" "${YAML_PATH}"
|
||||
else
|
||||
sed -i -e "/__REVIEWED_NEGATIVE_FEATURE_MARKER__/d" "${YAML_PATH}"
|
||||
fi
|
||||
|
||||
echo " Wrote training_parameters.yaml"
|
||||
rm -rf "${WORK_DIR}/trained_models/wakeword"
|
||||
|
||||
if [ "${IS_BLACKWELL}" = "true" ] && [ "${BLACKWELL_TF_MODE}" != "disabled" ]; then
|
||||
BLACKWELL_SETUP="${PROGDIR}/setup_blackwell_venv"
|
||||
BLACKWELL_PYTHON="${DATA_DIR}/.venv-blackwell/bin/python"
|
||||
|
||||
if [ -x "${BLACKWELL_SETUP}" ] && command -v python3.13 >/dev/null 2>&1; then
|
||||
echo "↪️ Preparing Blackwell-native TensorFlow training environment."
|
||||
if "${BLACKWELL_SETUP}" --data-dir="${DATA_DIR}"; then
|
||||
PYTHON_BIN="${BLACKWELL_PYTHON}"
|
||||
BLACKWELL_TF_ACTIVE="true"
|
||||
echo "✅ Blackwell TensorFlow training enabled: ${PYTHON_BIN}"
|
||||
else
|
||||
if [ "${BLACKWELL_TF_REQUIRED}" = "true" ]; then
|
||||
echo "❌ Blackwell TensorFlow setup failed and MWW_BLACKWELL_TF is required." >&2
|
||||
exit 1
|
||||
fi
|
||||
echo "⚠️ Blackwell TensorFlow setup failed; continuing with compatibility retries."
|
||||
fi
|
||||
else
|
||||
if [ "${BLACKWELL_TF_REQUIRED}" = "true" ]; then
|
||||
echo "❌ Blackwell TensorFlow was required, but python3.13/setup_blackwell_venv is unavailable." >&2
|
||||
exit 1
|
||||
fi
|
||||
echo "ℹ️ Blackwell TensorFlow image support not available; continuing with compatibility retries."
|
||||
fi
|
||||
fi
|
||||
|
||||
wake_word_filename="$(
|
||||
echo "${WAKE_WORD}" \
|
||||
| tr '[:upper:]' '[:lower:]' \
|
||||
@@ -234,11 +302,11 @@ TRAIN_ARGS=(
|
||||
--test_tflite_streaming_quantized 1
|
||||
--use_weights best_weights
|
||||
mixednet
|
||||
--pointwise_filters "64,64,64,64"
|
||||
--pointwise_filters "128,128,128,128"
|
||||
--repeat_in_block "1,1,1,1"
|
||||
--mixconv_kernel_sizes "[5], [7,11], [9,15], [23]"
|
||||
--residual_connection "0,0,0,0"
|
||||
--first_conv_filters 32
|
||||
--first_conv_filters 64
|
||||
--first_conv_kernel_size 5
|
||||
--stride 2
|
||||
)
|
||||
@@ -317,6 +385,8 @@ fi
|
||||
|
||||
TRAINING_DONE="false"
|
||||
|
||||
echo "🏋️ Starting model training and TFLite export (this is the longest stage)…"
|
||||
echo "🧠 Model quality: high_accuracy_plus"
|
||||
if run_attempt "Attempt 1/3: GPU training (default runtime profile)" ; then
|
||||
echo "✅ Training complete (GPU path)."
|
||||
TRAINING_DONE="true"
|
||||
@@ -386,12 +456,29 @@ if [ "${TRAINING_DONE}" != "true" ]; then
|
||||
fi
|
||||
|
||||
source_path="${WORK_DIR}/trained_models/wakeword/tflite_stream_state_internal_quant/stream_state_internal_quant.tflite"
|
||||
calibration_path="${WORK_DIR}/trained_models/wakeword/tflite_stream_state_internal_quant/detection_calibration.json"
|
||||
|
||||
if [ ! -f "${source_path}" ] ; then
|
||||
echo "Output model not found! Training didn't complete successfully. See ${TRAIN_LOG}"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "🎯 Calibrating detector settings for on-device use…"
|
||||
if "${PYTHON_BIN:-python}" "${PROGDIR}/calibrate_detector.py" \
|
||||
--training-config "${WORK_DIR}/trained_models/wakeword/training_config.yaml" \
|
||||
--model "${source_path}" \
|
||||
--output "${calibration_path}" \
|
||||
--target-faph "${MWW_CALIBRATION_TARGET_FAPH:-0.25}" \
|
||||
--recall-margin "${MWW_CALIBRATION_RECALL_MARGIN:-0.005}" \
|
||||
--window-sizes "${MWW_CALIBRATION_WINDOW_SIZES:-5,6,7}" \
|
||||
--cutoff-min "${MWW_CALIBRATION_CUTOFF_MIN:-0.95}" \
|
||||
--cutoff-max "${MWW_CALIBRATION_CUTOFF_MAX:-1.00}"; then
|
||||
echo "✅ Detector calibration complete."
|
||||
else
|
||||
echo "⚠️ Detector calibration failed; packaging with default detector settings."
|
||||
rm -f "${calibration_path}" || :
|
||||
fi
|
||||
|
||||
cp "${WORK_DIR}/trained_models/wakeword/model_summary.txt" "${OUTPUT_DIR}/logs/" || :
|
||||
cp -a "${WORK_DIR}/trained_models/wakeword/logs/train" "${OUTPUT_DIR}/logs/" || :
|
||||
cp -a "${WORK_DIR}/trained_models/wakeword/logs/validation" "${OUTPUT_DIR}/logs/" || :
|
||||
@@ -404,24 +491,93 @@ tflite_path="${OUTPUT_DIR}/${tflite_filename}"
|
||||
cp "${source_path}" "${tflite_path}"
|
||||
|
||||
json_path="${OUTPUT_DIR}/${wake_word_filename}.json"
|
||||
cat <<-EOF > "${json_path}"
|
||||
{
|
||||
export WAKE_WORD_TITLE LANGUAGE JSON_PATH="${json_path}" TFLITE_FILENAME="${tflite_filename}" CALIBRATION_PATH="${calibration_path}"
|
||||
echo "📦 Packaging final model artifacts…"
|
||||
"${PYTHON_BIN:-python}" - <<'PY'
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
json_path = Path(os.environ["JSON_PATH"])
|
||||
calibration_path = Path(os.environ.get("CALIBRATION_PATH", ""))
|
||||
language = (os.environ.get("LANGUAGE", "en") or "en").strip().lower()
|
||||
probability_cutoff = 0.97
|
||||
sliding_window_size = 6
|
||||
strict_min_close_miss_threshold = 0.68
|
||||
calibration = {}
|
||||
|
||||
if calibration_path.exists():
|
||||
try:
|
||||
calibration = json.loads(calibration_path.read_text(encoding="utf-8"))
|
||||
probability_cutoff = float(calibration.get("probability_cutoff", probability_cutoff))
|
||||
sliding_window_size = int(calibration.get("sliding_window_size", sliding_window_size))
|
||||
print(
|
||||
f"🎯 Using calibrated detector settings: "
|
||||
f"cutoff={probability_cutoff:.2f}, window={sliding_window_size}"
|
||||
)
|
||||
except Exception as exc:
|
||||
print(f"⚠️ Failed to read detector calibration ({exc}); using defaults.")
|
||||
|
||||
probability_cutoff = round(probability_cutoff, 3)
|
||||
sliding_window_size = max(1, min(10, int(sliding_window_size)))
|
||||
selected_metrics = calibration.get("selected_metrics") if isinstance(calibration.get("selected_metrics"), dict) else {}
|
||||
evaluation = calibration.get("evaluation") if isinstance(calibration.get("evaluation"), dict) else {}
|
||||
close_miss_threshold = max(
|
||||
0.01,
|
||||
min(0.99, round(max(strict_min_close_miss_threshold, probability_cutoff - 0.17), 3)),
|
||||
)
|
||||
|
||||
meta = {
|
||||
"type": "micro",
|
||||
"wake_word": "${WAKE_WORD_TITLE}",
|
||||
"wake_word": os.environ["WAKE_WORD_TITLE"],
|
||||
"label": os.environ["WAKE_WORD_TITLE"].replace("_", " ").title(),
|
||||
"author": "Tater Totterson",
|
||||
"website": "https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git",
|
||||
"model": "${tflite_filename}",
|
||||
"trained_languages": ["en"],
|
||||
"model": os.environ["TFLITE_FILENAME"],
|
||||
"trained_languages": [language],
|
||||
"version": 2,
|
||||
"model_format": "tflite_stream_state_internal_quant",
|
||||
"quantization": "int8",
|
||||
"sample_rate": 16000,
|
||||
"micro": {
|
||||
"probability_cutoff": 0.97,
|
||||
"sliding_window_size": 5,
|
||||
"probability_cutoff": probability_cutoff,
|
||||
"sliding_window_size": sliding_window_size,
|
||||
"feature_step_size": 10,
|
||||
"tensor_arena_size": 30000,
|
||||
"minimum_esphome_version": "2024.7.0"
|
||||
}
|
||||
"minimum_esphome_version": "2024.7.0",
|
||||
},
|
||||
"tater_native": {
|
||||
"format_version": 1,
|
||||
"wake_threshold": probability_cutoff,
|
||||
"wake_sliding_window": sliding_window_size,
|
||||
"close_miss_threshold": close_miss_threshold,
|
||||
"frontend": {
|
||||
"name": "tflm_microfrontend",
|
||||
"sample_rate": 16000,
|
||||
"feature_duration_ms": 30,
|
||||
"feature_step_ms": 10,
|
||||
"feature_size": 40,
|
||||
"input_feature_frames": 2,
|
||||
"lower_band_limit": 125.0,
|
||||
"upper_band_limit": 7500.0,
|
||||
},
|
||||
"recommended_for": ["tater-native-satellite", "voice-pe"],
|
||||
},
|
||||
"calibration": {
|
||||
"target_false_accepts_per_hour": calibration.get("target_false_accepts_per_hour"),
|
||||
"selected_false_accepts_per_hour_limit": calibration.get("selected_false_accepts_per_hour_limit"),
|
||||
"recall": selected_metrics.get("recall"),
|
||||
"false_accepts_per_hour": selected_metrics.get("false_accepts_per_hour"),
|
||||
"ambient_hours": selected_metrics.get("ambient_hours"),
|
||||
"positive_dataset": evaluation.get("positive_dataset"),
|
||||
"ambient_dataset": evaluation.get("ambient_dataset"),
|
||||
"positive_tracks": evaluation.get("positive_tracks"),
|
||||
"ambient_tracks": evaluation.get("ambient_tracks"),
|
||||
"generated_at": calibration.get("generated_at"),
|
||||
},
|
||||
}
|
||||
EOF
|
||||
json_path.write_text(json.dumps(meta, indent=4) + "\n", encoding="utf-8")
|
||||
PY
|
||||
|
||||
echo "Name: ${WAKE_WORD_TITLE}"
|
||||
echo "Model: ${tflite_path}"
|
||||
|
||||
@@ -6,7 +6,7 @@ ENV DEBIAN_FRONTEND=noninteractive
|
||||
# System deps
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
python3.12 python3.12-venv python3.12-dev python3-pip python-is-python3 \
|
||||
git wget curl unzip ca-certificates nano less \
|
||||
git wget curl unzip patch ninja-build ca-certificates nano less libgomp1 \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& mkdir -p /data
|
||||
|
||||
@@ -22,7 +22,7 @@ COPY --chown=root:root --chmod=0755 .bashrc /root/
|
||||
# Root-level entrypoints
|
||||
COPY --chown=root:root --chmod=0755 \
|
||||
train_wake_word \
|
||||
run_recorder.sh \
|
||||
run.sh \
|
||||
trainer_server.py \
|
||||
requirements.txt \
|
||||
/root/mww-scripts/
|
||||
@@ -37,4 +37,4 @@ RUN chmod -R a+x /root/mww-scripts/cli
|
||||
COPY --chown=root:root --chmod=0644 static/index.html /root/mww-scripts/static/index.html
|
||||
|
||||
# trainer server
|
||||
CMD ["/bin/bash", "-lc", "/root/mww-scripts/run_recorder.sh"]
|
||||
CMD ["/bin/bash", "-lc", "/root/mww-scripts/run.sh"]
|
||||
|
||||
54
dockerfile.blackwell
Normal file
54
dockerfile.blackwell
Normal file
@@ -0,0 +1,54 @@
|
||||
# RTX 50 / Blackwell image
|
||||
FROM nvidia/cuda:12.8.1-cudnn-devel-ubuntu24.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
ENV CUDA_HOME=/usr/local/cuda
|
||||
ENV PATH=/usr/local/cuda/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=/usr/local/cuda/lib64:${LD_LIBRARY_PATH}
|
||||
ENV MWW_BLACKWELL_IMAGE=1
|
||||
ENV MWW_BLACKWELL_TF=auto
|
||||
ENV MWW_BLACKWELL_TF_WHEEL_URL=https://github.com/chivitiH/tensorflow-blackwell-python313/releases/download/v2.22.0-selfbuilt/tensorflow-2.22.0.dev0+selfbuilt-cp313-cp313-linux_x86_64.whl
|
||||
|
||||
# System deps. Python 3.12 remains the main trainer/runtime venv, while
|
||||
# Python 3.13 is used only for the Blackwell TensorFlow training step.
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
software-properties-common ca-certificates curl git wget unzip patch \
|
||||
ninja-build nano less libgomp1 \
|
||||
&& add-apt-repository -y ppa:deadsnakes/ppa \
|
||||
&& apt-get update \
|
||||
&& apt-get install -y --no-install-recommends \
|
||||
python3.12 python3.12-venv python3.12-dev \
|
||||
python3.13 python3.13-venv python3.13-dev \
|
||||
python3-pip python-is-python3 \
|
||||
&& ldconfig \
|
||||
&& rm -rf /var/lib/apt/lists/* \
|
||||
&& mkdir -p /data
|
||||
|
||||
# Trainer UI port
|
||||
EXPOSE 8789
|
||||
|
||||
# Script root
|
||||
WORKDIR /root/mww-scripts
|
||||
|
||||
# Bash environment
|
||||
COPY --chown=root:root --chmod=0755 .bashrc /root/
|
||||
|
||||
# Root-level entrypoints
|
||||
COPY --chown=root:root --chmod=0755 \
|
||||
train_wake_word \
|
||||
run.sh \
|
||||
trainer_server.py \
|
||||
requirements.txt \
|
||||
/root/mww-scripts/
|
||||
|
||||
# CLI folder
|
||||
COPY --chown=root:root cli/ /root/mww-scripts/cli/
|
||||
|
||||
# Make all CLI scripts executable (avoids "Permission denied")
|
||||
RUN chmod -R a+x /root/mww-scripts/cli
|
||||
|
||||
# Static UI for trainer
|
||||
COPY --chown=root:root --chmod=0644 static/index.html /root/mww-scripts/static/index.html
|
||||
|
||||
# trainer server
|
||||
CMD ["/bin/bash", "-lc", "/root/mww-scripts/run.sh"]
|
||||
BIN
images/tater-repo-logo.png
Normal file
BIN
images/tater-repo-logo.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 590 KiB |
144
run.sh
Normal file
144
run.sh
Normal file
@@ -0,0 +1,144 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
ROOTDIR="$(dirname "$(realpath "$0")")"
|
||||
|
||||
# Training convention
|
||||
DATA_DIR="${DATA_DIR:-/data}"
|
||||
HOST="${REC_HOST:-0.0.0.0}"
|
||||
PORT="${REC_PORT:-8789}"
|
||||
|
||||
# Keep trainer UI deps separate from the training venv
|
||||
VENV_DIR="${DATA_DIR}/.recorder-venv"
|
||||
PY="${VENV_DIR}/bin/python"
|
||||
PIP="${PY} -m pip"
|
||||
PIN_FILE="${VENV_DIR}/.pinned_installed"
|
||||
|
||||
FASTAPI_VERSION="${REC_FASTAPI_VERSION:-0.115.6}"
|
||||
UVICORN_VERSION="${REC_UVICORN_VERSION:-0.30.6}"
|
||||
PY_MULTIPART_VERSION="${REC_PY_MULTIPART_VERSION:-0.0.9}"
|
||||
|
||||
echo "microWakeWord Trainer UI (Docker)"
|
||||
echo "-> ROOTDIR: ${ROOTDIR}"
|
||||
echo "-> DATA_DIR: ${DATA_DIR}"
|
||||
echo "-> URL: http://localhost:${PORT}/"
|
||||
|
||||
mkdir -p "${DATA_DIR}"
|
||||
|
||||
install_ui_deps() {
|
||||
${PIP} install \
|
||||
"fastapi==${FASTAPI_VERSION}" \
|
||||
"uvicorn[standard]==${UVICORN_VERSION}" \
|
||||
"python-multipart==${PY_MULTIPART_VERSION}" \
|
||||
"silero-vad>=5.0.0" \
|
||||
"numpy>=1.24.0" \
|
||||
"faster-whisper>=1.0.0" \
|
||||
"nvidia-cublas-cu12" \
|
||||
"nvidia-cudnn-cu12==9.*"
|
||||
}
|
||||
|
||||
# -----------------------------
|
||||
# Trainer UI venv (separate)
|
||||
# -----------------------------
|
||||
if [[ ! -x "${PY}" ]]; then
|
||||
echo "Creating trainer UI venv: ${VENV_DIR}"
|
||||
python3 -m venv "${VENV_DIR}"
|
||||
fi
|
||||
|
||||
# shellcheck disable=SC1091
|
||||
source "${VENV_DIR}/bin/activate"
|
||||
|
||||
if [[ ! -f "${PIN_FILE}" ]]; then
|
||||
echo "Installing pinned trainer UI deps"
|
||||
${PIP} install -U pip setuptools wheel
|
||||
install_ui_deps
|
||||
touch "${PIN_FILE}"
|
||||
else
|
||||
echo "Reusing existing trainer UI venv (no upgrades)"
|
||||
if ! "${PY}" - "${FASTAPI_VERSION}" "${UVICORN_VERSION}" "${PY_MULTIPART_VERSION}" <<'PY' >/dev/null 2>&1
|
||||
import importlib.metadata as md
|
||||
import sys
|
||||
|
||||
fastapi_version, uvicorn_version, multipart_version = sys.argv[1:4]
|
||||
|
||||
def version_tuple(value):
|
||||
parts = []
|
||||
for token in str(value).replace("-", ".").split("."):
|
||||
if token.isdigit():
|
||||
parts.append(int(token))
|
||||
else:
|
||||
digits = "".join(ch for ch in token if ch.isdigit())
|
||||
if digits:
|
||||
parts.append(int(digits))
|
||||
break
|
||||
return tuple(parts)
|
||||
|
||||
exact = {
|
||||
"fastapi": fastapi_version,
|
||||
"uvicorn": uvicorn_version,
|
||||
"python-multipart": multipart_version,
|
||||
}
|
||||
minimum = {
|
||||
"silero-vad": "5.0.0",
|
||||
"numpy": "1.24.0",
|
||||
"faster-whisper": "1.0.0",
|
||||
"nvidia-cudnn-cu12": "9.0.0",
|
||||
}
|
||||
present = (
|
||||
"torch",
|
||||
"nvidia-cublas-cu12",
|
||||
)
|
||||
|
||||
for package, expected in exact.items():
|
||||
if md.version(package) != expected:
|
||||
raise SystemExit(1)
|
||||
for package, minimum_version in minimum.items():
|
||||
if version_tuple(md.version(package)) < version_tuple(minimum_version):
|
||||
raise SystemExit(1)
|
||||
for package in present:
|
||||
md.version(package)
|
||||
PY
|
||||
then
|
||||
echo "UI dependencies missing or stale; installing recorder dependencies"
|
||||
install_ui_deps
|
||||
fi
|
||||
fi
|
||||
|
||||
# Faster Whisper/CTranslate2 loads these CUDA libraries before Python starts.
|
||||
# They live in the persistent UI venv so both Docker image variants can use GPU STT.
|
||||
WHISPER_CUDA_LIBRARY_PATH="$("${PY}" - <<'PY'
|
||||
import os
|
||||
|
||||
try:
|
||||
import nvidia.cublas.lib
|
||||
import nvidia.cudnn.lib
|
||||
except ImportError:
|
||||
print("")
|
||||
else:
|
||||
print(
|
||||
os.path.dirname(nvidia.cublas.lib.__file__)
|
||||
+ ":"
|
||||
+ os.path.dirname(nvidia.cudnn.lib.__file__)
|
||||
)
|
||||
PY
|
||||
)"
|
||||
if [[ -n "${WHISPER_CUDA_LIBRARY_PATH}" ]]; then
|
||||
export LD_LIBRARY_PATH="${WHISPER_CUDA_LIBRARY_PATH}${LD_LIBRARY_PATH:+:${LD_LIBRARY_PATH}}"
|
||||
fi
|
||||
# -----------------------------
|
||||
# Trainer server env
|
||||
# -----------------------------
|
||||
export DATA_DIR="${DATA_DIR}"
|
||||
export STATIC_DIR="${ROOTDIR}/static"
|
||||
export PERSONAL_DIR="${DATA_DIR}/personal_samples"
|
||||
export CAPTURED_DIR="${DATA_DIR}/captured_audio"
|
||||
export NEGATIVE_DIR="${DATA_DIR}/negative_samples"
|
||||
export TRAINED_WAKE_WORDS_DIR="${DATA_DIR}/trained_wake_words"
|
||||
|
||||
# IMPORTANT: leave training venv creation to /api/train inside trainer_server.py
|
||||
# but still set TRAIN_CMD so the server knows how to invoke training once ready
|
||||
export TRAIN_CMD="source '${DATA_DIR}/.venv/bin/activate' && train_wake_word --data-dir='${DATA_DIR}'"
|
||||
|
||||
echo "Launching uvicorn on ${HOST}:${PORT}"
|
||||
cd "${ROOTDIR}"
|
||||
exec "${VENV_DIR}/bin/uvicorn" trainer_server:app --host "${HOST}" --port "${PORT}"
|
||||
@@ -1,64 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
ROOTDIR="$(dirname "$(realpath "$0")")"
|
||||
|
||||
# Training convention
|
||||
DATA_DIR="${DATA_DIR:-/data}"
|
||||
HOST="${REC_HOST:-0.0.0.0}"
|
||||
PORT="${REC_PORT:-8888}"
|
||||
|
||||
# Keep recorder deps separate from training venv
|
||||
VENV_DIR="${DATA_DIR}/.recorder-venv"
|
||||
PY="${VENV_DIR}/bin/python"
|
||||
PIP="${PY} -m pip"
|
||||
PIN_FILE="${VENV_DIR}/.pinned_installed"
|
||||
|
||||
FASTAPI_VERSION="${REC_FASTAPI_VERSION:-0.115.6}"
|
||||
UVICORN_VERSION="${REC_UVICORN_VERSION:-0.30.6}"
|
||||
PY_MULTIPART_VERSION="${REC_PY_MULTIPART_VERSION:-0.0.9}"
|
||||
|
||||
echo "microWakeWord Trainer UI (Docker)"
|
||||
echo "-> ROOTDIR: ${ROOTDIR}"
|
||||
echo "-> DATA_DIR: ${DATA_DIR}"
|
||||
echo "-> URL: http://localhost:${PORT}/"
|
||||
|
||||
mkdir -p "${DATA_DIR}"
|
||||
|
||||
# -----------------------------
|
||||
# Trainer UI venv (separate)
|
||||
# -----------------------------
|
||||
if [[ ! -x "${PY}" ]]; then
|
||||
echo "Creating trainer UI venv: ${VENV_DIR}"
|
||||
python3 -m venv "${VENV_DIR}"
|
||||
fi
|
||||
|
||||
# shellcheck disable=SC1091
|
||||
source "${VENV_DIR}/bin/activate"
|
||||
|
||||
if [[ ! -f "${PIN_FILE}" ]]; then
|
||||
echo "Installing pinned trainer UI deps"
|
||||
${PIP} install -U pip setuptools wheel
|
||||
${PIP} install \
|
||||
"fastapi==${FASTAPI_VERSION}" \
|
||||
"uvicorn[standard]==${UVICORN_VERSION}" \
|
||||
"python-multipart==${PY_MULTIPART_VERSION}"
|
||||
touch "${PIN_FILE}"
|
||||
else
|
||||
echo "Reusing existing trainer UI venv (no upgrades)"
|
||||
fi
|
||||
|
||||
# -----------------------------
|
||||
# Trainer server env
|
||||
# -----------------------------
|
||||
export DATA_DIR="${DATA_DIR}"
|
||||
export STATIC_DIR="${ROOTDIR}/static"
|
||||
export PERSONAL_DIR="${DATA_DIR}/personal_samples"
|
||||
|
||||
# IMPORTANT: leave training venv creation to /api/train inside trainer_server.py
|
||||
# but still set TRAIN_CMD so the server knows how to invoke training once ready
|
||||
export TRAIN_CMD="source '${DATA_DIR}/.venv/bin/activate' && train_wake_word --data-dir='${DATA_DIR}'"
|
||||
|
||||
echo "Launching uvicorn on ${HOST}:${PORT}"
|
||||
cd "${ROOTDIR}"
|
||||
exec "${VENV_DIR}/bin/uvicorn" trainer_server:app --host "${HOST}" --port "${PORT}"
|
||||
2663
static/index.html
2663
static/index.html
File diff suppressed because it is too large
Load Diff
430
tests/test_auto_train.py
Normal file
430
tests/test_auto_train.py
Normal file
@@ -0,0 +1,430 @@
|
||||
import io
|
||||
import json
|
||||
import queue
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
import wave
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import trainer_server as trainer
|
||||
|
||||
|
||||
def silent_wav_bytes(duration_s: float = 0.25) -> bytes:
|
||||
output = io.BytesIO()
|
||||
with wave.open(output, "wb") as wav_file:
|
||||
wav_file.setnchannels(1)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(16000)
|
||||
wav_file.writeframes(b"\x00\x00" * int(16000 * duration_s))
|
||||
return output.getvalue()
|
||||
|
||||
|
||||
class AutoTrainTests(unittest.TestCase):
|
||||
def clear_review_queue(self):
|
||||
while True:
|
||||
try:
|
||||
trainer.AUTO_TRAIN_REVIEW_QUEUE.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
else:
|
||||
trainer.AUTO_TRAIN_REVIEW_QUEUE.task_done()
|
||||
trainer.AUTO_TRAIN_QUEUED_FILES.clear()
|
||||
|
||||
def setUp(self):
|
||||
self.clear_review_queue()
|
||||
self.tempdir = tempfile.TemporaryDirectory()
|
||||
root = Path(self.tempdir.name)
|
||||
self.original_paths = (
|
||||
trainer.CAPTURED_DIR,
|
||||
trainer.NEGATIVE_DIR,
|
||||
trainer.PERSONAL_DIR,
|
||||
trainer.AUTO_TRAIN_CONFIG_FILE,
|
||||
trainer.AUTO_TRAIN_STATE_FILE,
|
||||
)
|
||||
trainer.CAPTURED_DIR = root / "captured_audio"
|
||||
trainer.NEGATIVE_DIR = root / "negative_samples"
|
||||
trainer.PERSONAL_DIR = root / "personal_samples"
|
||||
trainer.AUTO_TRAIN_CONFIG_FILE = root / "auto_train_config.json"
|
||||
trainer.AUTO_TRAIN_STATE_FILE = root / "auto_train_state.json"
|
||||
for directory in (trainer.CAPTURED_DIR, trainer.NEGATIVE_DIR, trainer.PERSONAL_DIR):
|
||||
directory.mkdir(parents=True)
|
||||
|
||||
self.original_config = dict(trainer.AUTO_TRAIN_CONFIG)
|
||||
self.original_state = dict(trainer.AUTO_TRAIN_STATE)
|
||||
trainer.AUTO_TRAIN_CONFIG.clear()
|
||||
trainer.AUTO_TRAIN_CONFIG.update(
|
||||
trainer._normalize_auto_train_config(
|
||||
{
|
||||
"enabled": True,
|
||||
"wake_phrase": "hey tater",
|
||||
"language": "en",
|
||||
"tater_url": "http://127.0.0.1:8501",
|
||||
"stt_device": "auto",
|
||||
"stt_compute_type": "auto",
|
||||
}
|
||||
)
|
||||
)
|
||||
trainer.AUTO_TRAIN_STATE.clear()
|
||||
trainer.AUTO_TRAIN_STATE.update(trainer.AUTO_TRAIN_DEFAULT_STATE)
|
||||
|
||||
def tearDown(self):
|
||||
(
|
||||
trainer.CAPTURED_DIR,
|
||||
trainer.NEGATIVE_DIR,
|
||||
trainer.PERSONAL_DIR,
|
||||
trainer.AUTO_TRAIN_CONFIG_FILE,
|
||||
trainer.AUTO_TRAIN_STATE_FILE,
|
||||
) = self.original_paths
|
||||
trainer.AUTO_TRAIN_CONFIG.clear()
|
||||
trainer.AUTO_TRAIN_CONFIG.update(self.original_config)
|
||||
trainer.AUTO_TRAIN_STATE.clear()
|
||||
trainer.AUTO_TRAIN_STATE.update(self.original_state)
|
||||
self.clear_review_queue()
|
||||
self.tempdir.cleanup()
|
||||
|
||||
def add_capture(
|
||||
self,
|
||||
name: str = "wake.wav",
|
||||
wake_word: str = "hey_tater",
|
||||
event_type: str = "wake_detected",
|
||||
blocked_by_vad: bool = False,
|
||||
) -> Path:
|
||||
audio_path = trainer.CAPTURED_DIR / name
|
||||
audio_path.write_bytes(silent_wav_bytes())
|
||||
trainer._write_sidecar_json(
|
||||
audio_path,
|
||||
{
|
||||
"original_name": name,
|
||||
"wake_word": wake_word,
|
||||
"event_type": event_type,
|
||||
"blocked_by_vad": blocked_by_vad,
|
||||
"review_status": "pending",
|
||||
},
|
||||
)
|
||||
return audio_path
|
||||
|
||||
def test_phrase_matching_normalizes_case_punctuation_and_underscores(self):
|
||||
self.assertTrue(trainer._transcript_contains_wake_phrase("Okay, HEY TATER!", "hey_tater"))
|
||||
self.assertFalse(trainer._transcript_contains_wake_phrase("Turn on the television", "hey tater"))
|
||||
|
||||
def test_phrase_miss_moves_wake_trigger_to_negative_samples(self):
|
||||
self.add_capture()
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper", return_value="turn on the kitchen lights"):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertFalse((trainer.CAPTURED_DIR / "wake.wav").exists())
|
||||
negatives = list(trainer.NEGATIVE_DIR.glob("*.wav"))
|
||||
self.assertEqual(len(negatives), 1)
|
||||
metadata = trainer._load_sidecar_json(negatives[0])
|
||||
self.assertTrue(metadata["auto_negative"])
|
||||
self.assertEqual(metadata["review_status"], "auto_approved_negative")
|
||||
self.assertEqual(metadata["transcript"], "turn on the kitchen lights")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["pending_negative_count"], 1)
|
||||
|
||||
def test_matching_phrase_stays_in_manual_review_inbox(self):
|
||||
audio_path = self.add_capture()
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper", return_value="hey tater turn on the lights"):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertTrue(audio_path.exists())
|
||||
self.assertFalse(list(trainer.NEGATIVE_DIR.glob("*.wav")))
|
||||
metadata = trainer._load_sidecar_json(audio_path)
|
||||
self.assertEqual(metadata["auto_review_status"], "wake_phrase_detected")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["pending_negative_count"], 0)
|
||||
|
||||
def test_matching_phrase_is_deleted_when_cleanup_is_enabled(self):
|
||||
audio_path = self.add_capture()
|
||||
trainer.AUTO_TRAIN_CONFIG["delete_confirmed_wakes"] = True
|
||||
with patch.object(
|
||||
trainer,
|
||||
"_transcribe_capture_with_faster_whisper",
|
||||
return_value="hey tater turn on the lights",
|
||||
):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertFalse(audio_path.exists())
|
||||
self.assertFalse(audio_path.with_suffix(".json").exists())
|
||||
self.assertFalse(list(trainer.PERSONAL_DIR.glob("*.wav")))
|
||||
self.assertFalse(list(trainer.NEGATIVE_DIR.glob("*.wav")))
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_review_result"], "deleted_confirmed_wake")
|
||||
|
||||
def test_cleanup_processes_previously_confirmed_wake_without_retranscribing(self):
|
||||
audio_path = self.add_capture()
|
||||
metadata = trainer._load_sidecar_json(audio_path)
|
||||
metadata.update(
|
||||
{
|
||||
"auto_review_status": "wake_phrase_detected",
|
||||
"transcript": "hey tater",
|
||||
}
|
||||
)
|
||||
trainer._write_sidecar_json(audio_path, metadata)
|
||||
trainer.AUTO_TRAIN_CONFIG["delete_confirmed_wakes"] = True
|
||||
|
||||
self.assertEqual(trainer._queue_pending_auto_reviews(), 1)
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper") as transcribe:
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
transcribe.assert_not_called()
|
||||
self.assertFalse(audio_path.exists())
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_review_transcript"], "hey tater")
|
||||
|
||||
def test_close_miss_is_not_transcribed_by_default(self):
|
||||
audio_path = self.add_capture(event_type="close_miss")
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper") as transcribe:
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
transcribe.assert_not_called()
|
||||
self.assertTrue(audio_path.exists())
|
||||
self.assertFalse(trainer._load_sidecar_json(audio_path).get("auto_review_status"))
|
||||
|
||||
def test_existing_close_miss_is_queued_when_promotion_is_enabled(self):
|
||||
self.add_capture(event_type="close_miss")
|
||||
self.assertEqual(trainer._queue_pending_auto_reviews(), 0)
|
||||
|
||||
trainer.AUTO_TRAIN_CONFIG["promote_close_misses"] = True
|
||||
self.assertEqual(trainer._queue_pending_auto_reviews(), 1)
|
||||
|
||||
def test_close_miss_with_phrase_is_promoted_when_enabled(self):
|
||||
self.add_capture(event_type="close_miss")
|
||||
trainer.AUTO_TRAIN_CONFIG["promote_close_misses"] = True
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper", return_value="hey tater"):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertFalse((trainer.CAPTURED_DIR / "wake.wav").exists())
|
||||
positives = list(trainer.PERSONAL_DIR.glob("*.wav"))
|
||||
self.assertEqual(len(positives), 1)
|
||||
metadata = trainer._load_sidecar_json(positives[0])
|
||||
self.assertTrue(metadata["auto_positive"])
|
||||
self.assertEqual(metadata["review_status"], "auto_approved_personal")
|
||||
self.assertEqual(metadata["transcript"], "hey tater")
|
||||
self.assertFalse(list(trainer.NEGATIVE_DIR.glob("*.wav")))
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["pending_negative_count"], 0)
|
||||
|
||||
def test_close_miss_without_phrase_stays_in_inbox(self):
|
||||
audio_path = self.add_capture(event_type="close_miss")
|
||||
trainer.AUTO_TRAIN_CONFIG["promote_close_misses"] = True
|
||||
with patch.object(
|
||||
trainer,
|
||||
"_transcribe_capture_with_faster_whisper",
|
||||
return_value="turn on the lights",
|
||||
):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertTrue(audio_path.exists())
|
||||
self.assertFalse(list(trainer.PERSONAL_DIR.glob("*.wav")))
|
||||
self.assertFalse(list(trainer.NEGATIVE_DIR.glob("*.wav")))
|
||||
metadata = trainer._load_sidecar_json(audio_path)
|
||||
self.assertEqual(metadata["auto_review_status"], "close_miss_phrase_not_detected")
|
||||
|
||||
def test_vad_blocked_close_miss_is_never_transcribed(self):
|
||||
audio_path = self.add_capture(event_type="close_miss", blocked_by_vad=True)
|
||||
trainer.AUTO_TRAIN_CONFIG["promote_close_misses"] = True
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper") as transcribe:
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
transcribe.assert_not_called()
|
||||
self.assertTrue(audio_path.exists())
|
||||
self.assertFalse(trainer._load_sidecar_json(audio_path).get("auto_review_status"))
|
||||
|
||||
def test_capture_for_another_wake_word_is_not_transcribed(self):
|
||||
audio_path = self.add_capture(wake_word="computer")
|
||||
with patch.object(trainer, "_transcribe_capture_with_faster_whisper") as transcribe:
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
transcribe.assert_not_called()
|
||||
self.assertTrue(audio_path.exists())
|
||||
metadata = trainer._load_sidecar_json(audio_path)
|
||||
self.assertEqual(metadata["auto_review_status"], "different_wake_phrase")
|
||||
|
||||
def test_due_schedule_starts_training_after_minimum_negatives(self):
|
||||
trainer.AUTO_TRAIN_CONFIG["schedule_hours"] = 24
|
||||
trainer.AUTO_TRAIN_CONFIG["minimum_new_negatives"] = 3
|
||||
trainer.AUTO_TRAIN_STATE["pending_negative_count"] = 3
|
||||
trainer.AUTO_TRAIN_STATE["next_run_at"] = "2000-01-01T00:00:00+00:00"
|
||||
with patch.object(trainer, "_start_auto_training", return_value={"ok": True, "started": True}) as start:
|
||||
trainer._maybe_run_scheduled_auto_training()
|
||||
|
||||
start.assert_called_once_with()
|
||||
self.assertTrue(trainer.AUTO_TRAIN_STATE["next_run_at"])
|
||||
|
||||
def test_tater_notification_sets_new_word_globally_with_token(self):
|
||||
trainer.AUTO_TRAIN_CONFIG.update(
|
||||
{
|
||||
"notify_satellites": True,
|
||||
"tater_url": "http://127.0.0.1:8501",
|
||||
"tater_link_token": "secret-token",
|
||||
}
|
||||
)
|
||||
|
||||
class Response:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return b'{"push":{"count":4}}'
|
||||
|
||||
trained_word = {
|
||||
"key": "hey_tater",
|
||||
"wake_word": "Hey Tater",
|
||||
"json_url": "http://10.4.20.210:8789/api/trained_wake_words/hey_tater.json",
|
||||
}
|
||||
with (
|
||||
patch.object(trainer, "_advertised_base_url", return_value="http://10.4.20.210:8789"),
|
||||
patch.object(trainer, "_list_trained_wake_words", return_value=[trained_word]) as catalog,
|
||||
patch.object(trainer, "urlopen", return_value=Response()) as open_url,
|
||||
):
|
||||
result = trainer._notify_tater_satellites("hey_tater")
|
||||
|
||||
self.assertTrue(result["ok"])
|
||||
self.assertEqual(result["count"], 4)
|
||||
self.assertEqual(result["wake_word"], "Hey Tater")
|
||||
self.assertEqual(result["wake_word_url"], trained_word["json_url"])
|
||||
catalog.assert_called_once_with("http://10.4.20.210:8789")
|
||||
self.assertEqual(open_url.call_count, 1)
|
||||
request = open_url.call_args.args[0]
|
||||
self.assertEqual(request.full_url, "http://127.0.0.1:8501/api/tater/satellite/v1/trainer/wake-word")
|
||||
self.assertEqual(request.get_method(), "POST")
|
||||
self.assertEqual(request.get_header("X-tater-trainer-token"), "secret-token")
|
||||
self.assertEqual(
|
||||
json.loads(request.data),
|
||||
{
|
||||
"wake_word_name": "hey_tater",
|
||||
"wake_word_url": trained_word["json_url"],
|
||||
},
|
||||
)
|
||||
|
||||
def test_tater_notification_fails_when_trained_word_is_missing(self):
|
||||
trainer.AUTO_TRAIN_CONFIG["tater_link_token"] = "secret-token"
|
||||
with (
|
||||
patch.object(trainer, "_advertised_base_url", return_value="http://10.4.20.210:8789"),
|
||||
patch.object(trainer, "_list_trained_wake_words", return_value=[]),
|
||||
patch.object(trainer, "urlopen") as open_url,
|
||||
):
|
||||
result = trainer._notify_tater_satellites("missing_word")
|
||||
|
||||
self.assertFalse(result["ok"])
|
||||
self.assertIn("missing_word", result["error"])
|
||||
open_url.assert_not_called()
|
||||
|
||||
def test_tater_notification_requires_secure_link(self):
|
||||
trainer.AUTO_TRAIN_CONFIG["tater_link_token"] = ""
|
||||
with patch.object(trainer, "urlopen") as open_url:
|
||||
result = trainer._notify_tater_satellites("hey_tater")
|
||||
|
||||
self.assertFalse(result["ok"])
|
||||
self.assertIn("not linked", result["error"])
|
||||
open_url.assert_not_called()
|
||||
|
||||
def test_claim_tater_link_uses_tater_code_and_keeps_token_private(self):
|
||||
class Response:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def read(self, *_args):
|
||||
return json.dumps(
|
||||
{
|
||||
"ok": True,
|
||||
"token": "a" * 43,
|
||||
"tater_name": "Tater",
|
||||
"linked_at": "2026-07-24T12:00:00+00:00",
|
||||
}
|
||||
).encode("utf-8")
|
||||
|
||||
with (
|
||||
patch.object(trainer, "_advertised_base_url", return_value="http://10.4.20.210:8789"),
|
||||
patch.object(trainer, "urlopen", return_value=Response()) as open_url,
|
||||
):
|
||||
result = trainer._claim_tater_link("http://127.0.0.1:8501", "ABCD-EFGH")
|
||||
|
||||
self.assertTrue(result["linked"])
|
||||
self.assertEqual(trainer.AUTO_TRAIN_CONFIG["tater_link_token"], "a" * 43)
|
||||
self.assertNotIn("tater_link_token", trainer._public_auto_train_config())
|
||||
request = open_url.call_args.args[0]
|
||||
self.assertEqual(
|
||||
request.full_url,
|
||||
"http://127.0.0.1:8501/api/tater/satellite/v1/trainer/link/claim",
|
||||
)
|
||||
payload = json.loads(request.data)
|
||||
self.assertEqual(payload["pairing_code"], "ABCDEFGH")
|
||||
self.assertEqual(payload["publish_base_url"], "http://10.4.20.210:8789")
|
||||
self.assertTrue(payload["trainer_id"])
|
||||
|
||||
def test_advertised_url_uses_non_loopback_browser_host(self):
|
||||
request = SimpleNamespace(
|
||||
base_url="http://192.168.1.50:8789/",
|
||||
url=SimpleNamespace(hostname="192.168.1.50", scheme="http", port=8789),
|
||||
)
|
||||
self.assertEqual(trainer._advertised_base_url(request), "http://192.168.1.50:8789")
|
||||
|
||||
def test_advertised_url_replaces_localhost_with_discovered_lan_host(self):
|
||||
request = SimpleNamespace(
|
||||
base_url="http://127.0.0.1:8789/",
|
||||
url=SimpleNamespace(hostname="127.0.0.1", scheme="http", port=8789),
|
||||
)
|
||||
with patch.object(trainer, "_discover_lan_ipv4", return_value="192.168.1.60"):
|
||||
self.assertEqual(trainer._advertised_base_url(request), "http://192.168.1.60:8789")
|
||||
|
||||
def test_configured_public_url_takes_precedence(self):
|
||||
trainer.AUTO_TRAIN_CONFIG["advertised_base_url"] = "http://trainer.local:8789"
|
||||
request = SimpleNamespace(
|
||||
base_url="http://127.0.0.1:8789/",
|
||||
url=SimpleNamespace(hostname="127.0.0.1", scheme="http", port=8789),
|
||||
)
|
||||
self.assertEqual(trainer._advertised_base_url(request), "http://trainer.local:8789")
|
||||
|
||||
def test_faster_whisper_auto_runtime_prefers_cuda_and_float16(self):
|
||||
fake_ctranslate2 = SimpleNamespace(get_cuda_device_count=lambda: 1)
|
||||
with patch.dict(sys.modules, {"ctranslate2": fake_ctranslate2}):
|
||||
self.assertEqual(
|
||||
trainer._resolve_faster_whisper_runtime("auto", "auto"),
|
||||
("cuda", "float16"),
|
||||
)
|
||||
|
||||
def test_faster_whisper_auto_runtime_falls_back_to_cpu_int8(self):
|
||||
fake_ctranslate2 = SimpleNamespace(get_cuda_device_count=lambda: 0)
|
||||
with patch.dict(sys.modules, {"ctranslate2": fake_ctranslate2}):
|
||||
self.assertEqual(
|
||||
trainer._resolve_faster_whisper_runtime("auto", "auto"),
|
||||
("cpu", "int8"),
|
||||
)
|
||||
|
||||
def test_faster_whisper_transcription_joins_segments_and_records_runtime(self):
|
||||
fake_model = SimpleNamespace()
|
||||
fake_model.transcribe = Mock(
|
||||
return_value=(
|
||||
iter([SimpleNamespace(text=" turn on "), SimpleNamespace(text="the lights ")]),
|
||||
SimpleNamespace(),
|
||||
)
|
||||
)
|
||||
with (
|
||||
patch.object(trainer, "_resolve_faster_whisper_runtime", return_value=("cuda", "float16")),
|
||||
patch.object(trainer, "_load_faster_whisper_model", return_value=fake_model),
|
||||
):
|
||||
transcript = trainer._transcribe_capture_with_faster_whisper(
|
||||
Path("wake.wav"),
|
||||
model="small.en",
|
||||
language="en",
|
||||
)
|
||||
|
||||
self.assertEqual(transcript, "turn on the lights")
|
||||
fake_model.transcribe.assert_called_once_with(
|
||||
"wake.wav",
|
||||
language="en",
|
||||
beam_size=1,
|
||||
condition_on_previous_text=False,
|
||||
)
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_stt_device"], "cuda")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_stt_compute_type"], "float16")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
88
tests/test_calibrate_detector.py
Normal file
88
tests/test_calibrate_detector.py
Normal file
@@ -0,0 +1,88 @@
|
||||
import importlib.util
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
SCRIPT_PATH = (
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "cli"
|
||||
/ "calibrate_detector.py"
|
||||
)
|
||||
SPEC = importlib.util.spec_from_file_location("calibrate_detector", SCRIPT_PATH)
|
||||
calibrate_detector = importlib.util.module_from_spec(SPEC)
|
||||
assert SPEC.loader is not None
|
||||
SPEC.loader.exec_module(calibrate_detector)
|
||||
|
||||
|
||||
def candidate(cutoff, window, recall, false_accepts_per_hour):
|
||||
return {
|
||||
"probability_cutoff": cutoff,
|
||||
"sliding_window_size": window,
|
||||
"recall": recall,
|
||||
"false_accepts_per_hour": false_accepts_per_hour,
|
||||
}
|
||||
|
||||
|
||||
class CalibrationSelectionTests(unittest.TestCase):
|
||||
def test_defaults_are_conservative(self):
|
||||
self.assertEqual(calibrate_detector.DEFAULT_WINDOW_SIZES, [5, 6, 7])
|
||||
self.assertEqual(calibrate_detector.DEFAULT_CUTOFF_MIN, 0.95)
|
||||
self.assertEqual(calibrate_detector.DEFAULT_RECALL_MARGIN, 0.005)
|
||||
|
||||
def test_prefers_zero_false_accepts_within_recall_margin(self):
|
||||
candidates = [
|
||||
candidate(0.95, 5, 0.99894, 0.103408),
|
||||
candidate(0.95, 6, 0.99744, 0.0),
|
||||
candidate(0.95, 7, 0.99554, 0.0),
|
||||
]
|
||||
|
||||
best, selected_limit = calibrate_detector._select_best_candidate(
|
||||
candidates,
|
||||
target_faph=0.25,
|
||||
recall_margin=0.005,
|
||||
)
|
||||
|
||||
self.assertEqual(best["sliding_window_size"], 6)
|
||||
self.assertEqual(best["false_accepts_per_hour"], 0.0)
|
||||
self.assertEqual(selected_limit, 0.25)
|
||||
|
||||
def test_does_not_trade_away_recall_beyond_margin(self):
|
||||
candidates = [
|
||||
candidate(0.95, 5, 0.99, 0.1),
|
||||
candidate(0.99, 6, 0.90, 0.0),
|
||||
]
|
||||
|
||||
best, _ = calibrate_detector._select_best_candidate(
|
||||
candidates,
|
||||
target_faph=0.25,
|
||||
recall_margin=0.005,
|
||||
)
|
||||
|
||||
self.assertEqual(best["sliding_window_size"], 5)
|
||||
|
||||
def test_uses_strictest_available_false_accept_tier(self):
|
||||
candidates = [
|
||||
candidate(0.95, 5, 0.99, 0.6),
|
||||
candidate(0.99, 6, 0.99, 1.5),
|
||||
]
|
||||
|
||||
best, selected_limit = calibrate_detector._select_best_candidate(
|
||||
candidates,
|
||||
target_faph=0.25,
|
||||
recall_margin=0.005,
|
||||
)
|
||||
|
||||
self.assertEqual(best["false_accepts_per_hour"], 0.6)
|
||||
self.assertEqual(selected_limit, 0.75)
|
||||
|
||||
def test_rejects_negative_recall_margin(self):
|
||||
with self.assertRaises(ValueError):
|
||||
calibrate_detector._select_best_candidate(
|
||||
[candidate(0.95, 6, 0.99, 0.0)],
|
||||
target_faph=0.25,
|
||||
recall_margin=-0.001,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -150,6 +150,7 @@ if ${CLEANUP_WORK_DIR} ; then
|
||||
"${DATA_DIR}/work/wake_word_samples" \
|
||||
"${DATA_DIR}/work/wake_word_samples_augmented" \
|
||||
"${DATA_DIR}/work/personal_augmented_features" \
|
||||
"${DATA_DIR}/work/reviewed_negative_features" \
|
||||
"${DATA_DIR}/work/last_wake_word" || :
|
||||
fi
|
||||
|
||||
|
||||
2358
trainer_server.py
2358
trainer_server.py
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user