mirror of
https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git
synced 2026-08-12 07:55:33 -06:00
Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7d77f71dc3 | ||
|
|
3d341d0617 | ||
|
|
a1b22200e0 | ||
|
|
89260f1f14 | ||
|
|
0140dfb56f | ||
|
|
1fc7d80bae | ||
|
|
31a6388da4 | ||
|
|
85c2d6334b |
88
.github/workflows/docker-publish.yml
vendored
88
.github/workflows/docker-publish.yml
vendored
@@ -7,7 +7,7 @@ on:
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
contents: write
|
||||
packages: write
|
||||
|
||||
concurrency:
|
||||
@@ -26,6 +26,18 @@ jobs:
|
||||
- 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
|
||||
|
||||
@@ -41,6 +53,8 @@ jobs:
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}
|
||||
flavor: |
|
||||
latest=false
|
||||
tags: |
|
||||
type=raw,value=latest
|
||||
type=ref,event=tag
|
||||
@@ -56,3 +70,75 @@ jobs:
|
||||
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
|
||||
|
||||
140
README.md
140
README.md
@@ -7,7 +7,7 @@
|
||||
<a href="https://taterassistant.com">taterassistant.com</a>
|
||||
</h3>
|
||||
|
||||
Train custom microWakeWord models in Docker with NVIDIA/CUDA acceleration, generated Piper samples, device-captured samples, reviewed false-wake negatives, live training logs, and prebuilt Tater firmware flashing.
|
||||
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.
|
||||
|
||||
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.
|
||||
|
||||
@@ -22,9 +22,24 @@ docker pull ghcr.io/tatertotterson/microwakeword:latest
|
||||
Tagged releases also publish matching immutable image tags:
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v5
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v12
|
||||
```
|
||||
|
||||
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:v12-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
|
||||
@@ -38,16 +53,19 @@ docker run -d \
|
||||
ghcr.io/tatertotterson/microwakeword:latest
|
||||
```
|
||||
|
||||
Use a version tag such as `ghcr.io/tatertotterson/microwakeword:v5` when you want to pin a known release instead of tracking `latest`.
|
||||
Use a version tag such as `ghcr.io/tatertotterson/microwakeword:v12` 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:v12-blackwell`
|
||||
in the same `docker run` command.
|
||||
|
||||
The flags:
|
||||
|
||||
- `--gpus all` enables GPU acceleration.
|
||||
- `--network host` lets the container receive mDNS/zeroconf traffic for ESPHome auto-detect.
|
||||
- `--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 firmware caches.
|
||||
- `-v $(pwd):/data` persists models, downloaded voices, datasets, samples, and generated wake-word artifacts.
|
||||
|
||||
Host networking is recommended for the Firmware tab's mDNS device discovery. Manual IP flashing and captured-audio uploads can still work without host networking if the trainer port is reachable, but auto-detect may not see devices from Docker bridge networking.
|
||||
If you do not use host networking, publish the trainer port and make sure satellites can reach it from your LAN.
|
||||
|
||||
Open:
|
||||
|
||||
@@ -55,31 +73,38 @@ Open:
|
||||
http://localhost:8789
|
||||
```
|
||||
|
||||
If you change `REC_PORT`, open that port instead and use the same port in the ESPHome `Trainer App URL`.
|
||||
If you change `REC_PORT`, open that port instead and use the same port in the satellite `Trainer App URL`.
|
||||
|
||||
---
|
||||
|
||||
## What The UI Does
|
||||
|
||||
- `Trainer` starts a wake-word session, shows positive/negative sample counts, and launches training.
|
||||
- `Captured Audio` reviews clips sent by ESPHome sats, including wake hits, close misses, and false wakes.
|
||||
- `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.
|
||||
- `Firmware` pulls verified prebuilt Tater firmware images from GitHub and flashes supported satellites over OTA.
|
||||
- Popup consoles show colorized training and firmware logs while long-running jobs are active.
|
||||
- `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.
|
||||
|
||||
---
|
||||
|
||||
## Captured Audio Workflow
|
||||
|
||||
To collect samples from a sat, flash it with the Tater firmware from [TaterTotterson/microWakeWords](https://github.com/TaterTotterson/microWakeWords). The `Firmware` tab can pull verified prebuilt OTA images from that repo for fast firmware updates.
|
||||
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.
|
||||
|
||||
After flashing, the device exposes ESPHome entities for capture setup:
|
||||
For Tater Native satellites, enable trainer feedback in Tater:
|
||||
|
||||
- `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`.
|
||||
|
||||
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`.
|
||||
|
||||
ESPHome devices can send raw captured audio to:
|
||||
Satellites send raw captured audio to:
|
||||
|
||||
```text
|
||||
/api/upload_captured_audio_raw
|
||||
@@ -140,6 +165,27 @@ Starting a new session does not clear samples. Use the clear buttons in `Samples
|
||||
|
||||
---
|
||||
|
||||
## Auto Training
|
||||
|
||||
`Auto Training` is an opt-in false-positive 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 positive review.
|
||||
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, close misses, VAD-blocked captures, and captures for another wake word stay out of the automatic negative path.
|
||||
|
||||
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 publishes the replacement model at the same wake-word URL and can call Tater's native satellite settings API to make connected satellites pull it again. This refresh uses the existing Tater Native update path, 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. The optional API token is stored in `/data/auto_train_config.json` with owner-only permissions.
|
||||
|
||||
---
|
||||
|
||||
## Training Flow
|
||||
|
||||
1. Enter the wake phrase in `Trainer`.
|
||||
@@ -153,6 +199,8 @@ Personal samples are optional. Training can run with zero personal samples after
|
||||
|
||||
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
|
||||
@@ -183,20 +231,17 @@ After those assets are prepared, later runs reuse the local copies unless the mo
|
||||
|
||||
---
|
||||
|
||||
## Firmware Flashing
|
||||
## Trained Wake Words
|
||||
|
||||
The `Firmware` tab flashes prebuilt Tater firmware for supported ESPHome satellites.
|
||||
The `Wake Words` tab lists locally trained wake-word packages from `/data/trained_wake_words/`.
|
||||
|
||||
- Downloads the latest prebuilt firmware manifest plus OTA and USB factory images from `TaterTotterson/microWakeWords`.
|
||||
- Verifies downloaded images by size and SHA before upload.
|
||||
- Auto-detects ESPHome devices with mDNS when the container is running with host networking.
|
||||
- Allows manual IP or hostname entry if discovery does not find the device.
|
||||
- Saves the selected OTA target for each firmware family.
|
||||
- Flashes the prebuilt factory image over Browser USB for first installs or recovery when opened in Chrome or Edge.
|
||||
- Lists locally trained wake words from `/data/trained_wake_words/` for live model switching.
|
||||
- Streams download, verification, and OTA upload progress in a colorized firmware console.
|
||||
- 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.
|
||||
|
||||
You usually only flash for firmware updates. New satellites, or devices older than Tater firmware `3.0.3`, need one USB flash first before OTA updates and live wake-word switching are available.
|
||||
Use the main Tater app for satellite firmware updates and USB flashing.
|
||||
|
||||
---
|
||||
|
||||
@@ -209,14 +254,51 @@ Successful runs produce timestamped training output folders such as:
|
||||
/data/output/<timestamp>-<wake_word>-<samples>-<steps>/<wake_word>.json
|
||||
```
|
||||
|
||||
The trainer also syncs firmware-ready artifacts into:
|
||||
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 firmware tab uses `/data/trained_wake_words/` to populate the wake-word dropdown.
|
||||
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`.
|
||||
|
||||
---
|
||||
|
||||
@@ -233,7 +315,7 @@ That removes:
|
||||
- cached datasets
|
||||
- training environments
|
||||
- trained models
|
||||
- downloaded firmware images
|
||||
- Auto Training settings, state, transcripts, and cached Faster Whisper models
|
||||
|
||||
---
|
||||
|
||||
@@ -241,9 +323,10 @@ That removes:
|
||||
|
||||
- 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`.
|
||||
- Firmware capture settings live on the ESPHome device and can be toggled from the device entities after flashing.
|
||||
- Trainer capture settings live in Tater for Tater Native satellites, and on device entities for older ESPHome satellites.
|
||||
|
||||
---
|
||||
|
||||
@@ -253,3 +336,4 @@ 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
|
||||
|
||||
5
WHATS_NEW.md
Normal file
5
WHATS_NEW.md
Normal file
@@ -0,0 +1,5 @@
|
||||
- Added opt-in Auto Training for false-positive wake triggers, using Faster Whisper with automatic CUDA/float16 selection and CPU/int8 fallback.
|
||||
- Wake triggers whose transcripts do not contain the configured phrase can now become hard negatives automatically; close misses, empty transcripts, and phrase matches remain available for manual review.
|
||||
- Added scheduled retraining with a minimum-new-negatives threshold and automatic Tater Native satellite refresh after a successful model build.
|
||||
- Wake-word download links now advertise a LAN-reachable trainer URL instead of `127.0.0.1`.
|
||||
- Tightened detector calibration defaults to favor fewer ambient false accepts while preserving candidates within 0.5 percentage points of the best recall.
|
||||
@@ -9,24 +9,22 @@ import math
|
||||
import os
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Iterable, Sequence
|
||||
from typing import Any, Iterable, Sequence
|
||||
|
||||
import numpy as np
|
||||
import yaml
|
||||
|
||||
from microwakeword.data import FeatureHandler
|
||||
from microwakeword.inference import Model
|
||||
|
||||
|
||||
DEFAULT_WINDOW_SIZES = [3, 4, 5, 6, 7]
|
||||
DEFAULT_TARGET_FAPH = float(os.environ.get("MWW_CALIBRATION_TARGET_FAPH", "1.0"))
|
||||
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.00"))
|
||||
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:
|
||||
@@ -65,6 +63,15 @@ def parse_args() -> argparse.Namespace:
|
||||
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,
|
||||
@@ -159,7 +166,13 @@ def _compute_false_accepts_per_hour(
|
||||
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),
|
||||
@@ -172,13 +185,27 @@ def _select_best_candidate(
|
||||
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(
|
||||
candidates,
|
||||
viable_candidates,
|
||||
key=lambda candidate: (
|
||||
tier(candidate),
|
||||
-candidate["recall"],
|
||||
candidate["false_accepts_per_hour"],
|
||||
abs(candidate["sliding_window_size"] - 5),
|
||||
-candidate["recall"],
|
||||
abs(candidate["sliding_window_size"] - PREFERRED_WINDOW_SIZE),
|
||||
-candidate["probability_cutoff"],
|
||||
),
|
||||
)
|
||||
@@ -195,7 +222,7 @@ def _load_config(config_path: Path) -> dict:
|
||||
|
||||
|
||||
def _load_eval_sets(
|
||||
handler: FeatureHandler,
|
||||
handler: Any,
|
||||
config: dict,
|
||||
) -> tuple[str, str, list[np.ndarray], list[np.ndarray]]:
|
||||
for positive_mode, ambient_mode in (
|
||||
@@ -228,7 +255,7 @@ def _load_eval_sets(
|
||||
|
||||
|
||||
def _predict_tracks(
|
||||
model: Model,
|
||||
model: Any,
|
||||
tracks: Sequence[np.ndarray],
|
||||
label: str,
|
||||
) -> list[np.ndarray]:
|
||||
@@ -244,8 +271,13 @@ def _predict_tracks(
|
||||
|
||||
|
||||
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:
|
||||
@@ -276,6 +308,10 @@ def main() -> int:
|
||||
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", {})
|
||||
@@ -338,7 +374,11 @@ def main() -> int:
|
||||
candidates.append(candidate)
|
||||
window_candidates.append(candidate)
|
||||
|
||||
best_window, _ = _select_best_candidate(window_candidates, args.target_faph)
|
||||
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%}; "
|
||||
@@ -350,7 +390,11 @@ def main() -> int:
|
||||
)
|
||||
)
|
||||
|
||||
best, selected_limit = _select_best_candidate(candidates, args.target_faph)
|
||||
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; "
|
||||
@@ -390,6 +434,8 @@ def main() -> int:
|
||||
"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(),
|
||||
|
||||
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}"
|
||||
@@ -45,6 +45,7 @@ download_hf_mit_rirs() {
|
||||
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
|
||||
@@ -67,24 +68,57 @@ files = sorted(
|
||||
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
|
||||
for rel in files:
|
||||
|
||||
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}"
|
||||
with urllib.request.urlopen(url, timeout=60) as response:
|
||||
target.write_bytes(response.read())
|
||||
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)")
|
||||
print(f" MIT environmental RIR files available: {len(files)}")
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -238,6 +253,32 @@ 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:]' \
|
||||
@@ -261,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
|
||||
)
|
||||
@@ -345,6 +386,7 @@ 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"
|
||||
@@ -425,7 +467,12 @@ 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}"; then
|
||||
--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."
|
||||
@@ -455,7 +502,9 @@ 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 = 5
|
||||
sliding_window_size = 6
|
||||
strict_min_close_miss_threshold = 0.68
|
||||
calibration = {}
|
||||
|
||||
if calibration_path.exists():
|
||||
try:
|
||||
@@ -469,21 +518,63 @@ if calibration_path.exists():
|
||||
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": 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": os.environ["TFLITE_FILENAME"],
|
||||
"trained_languages": [language],
|
||||
"version": 2,
|
||||
"model_format": "tflite_stream_state_internal_quant",
|
||||
"quantization": "int8",
|
||||
"sample_rate": 16000,
|
||||
"micro": {
|
||||
"probability_cutoff": round(probability_cutoff, 2),
|
||||
"probability_cutoff": probability_cutoff,
|
||||
"sliding_window_size": sliding_window_size,
|
||||
"feature_step_size": 10,
|
||||
"tensor_arena_size": 30000,
|
||||
"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"),
|
||||
},
|
||||
}
|
||||
json_path.write_text(json.dumps(meta, indent=4) + "\n", encoding="utf-8")
|
||||
PY
|
||||
|
||||
@@ -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 patch ninja-build 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
|
||||
|
||||
|
||||
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"]
|
||||
36
run.sh
36
run.sh
@@ -30,9 +30,11 @@ install_ui_deps() {
|
||||
"fastapi==${FASTAPI_VERSION}" \
|
||||
"uvicorn[standard]==${UVICORN_VERSION}" \
|
||||
"python-multipart==${PY_MULTIPART_VERSION}" \
|
||||
"zeroconf>=0.132.2" \
|
||||
"silero-vad>=5.0.0" \
|
||||
"numpy>=1.24.0"
|
||||
"numpy>=1.24.0" \
|
||||
"faster-whisper>=1.0.0" \
|
||||
"nvidia-cublas-cu12" \
|
||||
"nvidia-cudnn-cu12==9.*"
|
||||
}
|
||||
|
||||
# -----------------------------
|
||||
@@ -79,9 +81,13 @@ exact = {
|
||||
minimum = {
|
||||
"silero-vad": "5.0.0",
|
||||
"numpy": "1.24.0",
|
||||
"zeroconf": "0.132.2",
|
||||
"faster-whisper": "1.0.0",
|
||||
"nvidia-cudnn-cu12": "9.0.0",
|
||||
}
|
||||
present = ("torch",)
|
||||
present = (
|
||||
"torch",
|
||||
"nvidia-cublas-cu12",
|
||||
)
|
||||
|
||||
for package, expected in exact.items():
|
||||
if md.version(package) != expected:
|
||||
@@ -97,6 +103,28 @@ PY
|
||||
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
|
||||
# -----------------------------
|
||||
|
||||
1853
static/index.html
1853
static/index.html
File diff suppressed because it is too large
Load Diff
238
tests/test_auto_train.py
Normal file
238
tests/test_auto_train.py
Normal file
@@ -0,0 +1,238 @@
|
||||
import io
|
||||
import json
|
||||
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 setUp(self):
|
||||
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.tempdir.cleanup()
|
||||
|
||||
def add_capture(self, name: str = "wake.wav", wake_word: str = "hey_tater") -> 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": "wake_detected",
|
||||
"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_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_refresh_repushes_settings_with_selector_and_token(self):
|
||||
trainer.AUTO_TRAIN_CONFIG.update(
|
||||
{
|
||||
"notify_satellites": True,
|
||||
"tater_url": "http://127.0.0.1:8501",
|
||||
"tater_selector": "kitchen-sat",
|
||||
"tater_api_token": "secret-token",
|
||||
}
|
||||
)
|
||||
|
||||
class Response:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return b'{"push":{"count":2}}'
|
||||
|
||||
with patch.object(trainer, "urlopen", return_value=Response()) as open_url:
|
||||
result = trainer._notify_tater_satellites()
|
||||
|
||||
self.assertTrue(result["ok"])
|
||||
self.assertEqual(result["count"], 2)
|
||||
request = open_url.call_args.args[0]
|
||||
self.assertEqual(request.full_url, "http://127.0.0.1:8501/api/tater/satellite/v1/settings")
|
||||
self.assertEqual(request.get_header("X-tater-token"), "secret-token")
|
||||
self.assertEqual(json.loads(request.data), {"selector": "kitchen-sat", "settings": {}})
|
||||
|
||||
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()
|
||||
2591
trainer_server.py
2591
trainer_server.py
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user