mirror of
https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git
synced 2026-08-12 16:05:34 -06:00
Compare commits
9 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2b1320f1f3 | ||
|
|
19ee63a65b | ||
|
|
518df63161 | ||
|
|
6ee228e8d3 | ||
|
|
2eee70cb34 | ||
|
|
426e4ec83f | ||
|
|
c474deb8b5 | ||
|
|
931694b711 | ||
|
|
5554b2eb5e |
29
README.md
29
README.md
@@ -22,7 +22,7 @@ docker pull ghcr.io/tatertotterson/microwakeword:latest
|
||||
Tagged releases also publish matching immutable image tags:
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v12
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v17
|
||||
```
|
||||
|
||||
The release tag must match `VERSION`. Update `WHATS_NEW.md` before tagging; the Docker workflow prepends it to GitHub's automatically generated release notes.
|
||||
@@ -32,7 +32,7 @@ Python 3.13 TensorFlow build for `sm_120`:
|
||||
|
||||
```bash
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:blackwell
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v12-blackwell
|
||||
docker pull ghcr.io/tatertotterson/microwakeword:v17-blackwell
|
||||
```
|
||||
|
||||
Use the Blackwell image only for RTX 50-series cards. It includes the
|
||||
@@ -53,9 +53,9 @@ docker run -d \
|
||||
ghcr.io/tatertotterson/microwakeword:latest
|
||||
```
|
||||
|
||||
Use a version tag such as `ghcr.io/tatertotterson/microwakeword:v12` when you want to pin a known release instead of tracking `latest`.
|
||||
Use a version tag such as `ghcr.io/tatertotterson/microwakeword:v17` 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`
|
||||
or a pinned tag such as `ghcr.io/tatertotterson/microwakeword:v17-blackwell`
|
||||
in the same `docker run` command.
|
||||
|
||||
The flags:
|
||||
@@ -167,22 +167,29 @@ 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.
|
||||
`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 positive review.
|
||||
1. The selected local STT engine transcribes the audio.
|
||||
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, close misses, VAD-blocked captures, and captures for another wake word stay out of the automatic negative path.
|
||||
4. Empty transcripts, 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/`.
|
||||
Two optional cleanup rules are available:
|
||||
|
||||
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.
|
||||
- `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.
|
||||
|
||||
Auto Training exposes only an engine selector. Faster Whisper is the recommended default and uses CUDA with `float16` when CTranslate2 can see an NVIDIA GPU, with a CPU `int8` fallback. The trainer manages `small.en` for English and `small` for other languages. Parakeet ONNX uses the managed INT8 `nemo-parakeet-tdt-0.6b-v3` model with CUDA and CPU fallback. 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. The optional API token is stored in `/data/auto_train_config.json` with owner-only permissions.
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
- 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.
|
||||
- Improved automatic review accuracy for short wake phrases that STT initially hears as similar-sounding words.
|
||||
- Added a conservative Faster Whisper confirmation pass that uses the currently configured wake phrase only when the unbiased transcript is already phonetically close.
|
||||
- Kept unconfirmed close transcripts in the manual review inbox instead of allowing them to become harmful negative training samples.
|
||||
- Added visible guided-transcript and review-reason details, plus retry support for ambiguous clips through Review Now.
|
||||
|
||||
@@ -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}" || :
|
||||
|
||||
47
run.sh
47
run.sh
@@ -33,8 +33,11 @@ install_ui_deps() {
|
||||
"silero-vad>=5.0.0" \
|
||||
"numpy>=1.24.0" \
|
||||
"faster-whisper>=1.0.0" \
|
||||
"onnx-asr[hub]>=0.12.0" \
|
||||
"nvidia-cublas-cu12" \
|
||||
"nvidia-cudnn-cu12==9.*"
|
||||
${PIP} uninstall -y onnxruntime
|
||||
${PIP} install "onnxruntime-gpu[cuda,cudnn]<1.27"
|
||||
}
|
||||
|
||||
# -----------------------------
|
||||
@@ -82,11 +85,13 @@ minimum = {
|
||||
"silero-vad": "5.0.0",
|
||||
"numpy": "1.24.0",
|
||||
"faster-whisper": "1.0.0",
|
||||
"onnx-asr": "0.12.0",
|
||||
"nvidia-cudnn-cu12": "9.0.0",
|
||||
}
|
||||
present = (
|
||||
"torch",
|
||||
"nvidia-cublas-cu12",
|
||||
"onnxruntime-gpu",
|
||||
)
|
||||
|
||||
for package, expected in exact.items():
|
||||
@@ -97,6 +102,10 @@ for package, minimum_version in minimum.items():
|
||||
raise SystemExit(1)
|
||||
for package in present:
|
||||
md.version(package)
|
||||
|
||||
import onnxruntime as ort
|
||||
if "CUDAExecutionProvider" not in ort.get_available_providers():
|
||||
raise SystemExit(1)
|
||||
PY
|
||||
then
|
||||
echo "UI dependencies missing or stale; installing recorder dependencies"
|
||||
@@ -107,19 +116,33 @@ 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
|
||||
from importlib.util import find_spec
|
||||
from pathlib import Path
|
||||
|
||||
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__)
|
||||
)
|
||||
|
||||
def package_directory(name):
|
||||
try:
|
||||
spec = find_spec(name)
|
||||
except (ImportError, AttributeError, ValueError):
|
||||
return ""
|
||||
if spec is None:
|
||||
return ""
|
||||
|
||||
for location in spec.submodule_search_locations or ():
|
||||
if location:
|
||||
return str(Path(location).resolve())
|
||||
|
||||
origin = spec.origin
|
||||
if origin and origin not in {"built-in", "frozen"}:
|
||||
return str(Path(origin).resolve().parent)
|
||||
return ""
|
||||
|
||||
|
||||
paths = [
|
||||
package_directory("nvidia.cublas.lib"),
|
||||
package_directory("nvidia.cudnn.lib"),
|
||||
]
|
||||
print(":".join(dict.fromkeys(path for path in paths if path)))
|
||||
PY
|
||||
)"
|
||||
if [[ -n "${WHISPER_CUDA_LIBRARY_PATH}" ]]; then
|
||||
|
||||
@@ -1084,6 +1084,43 @@
|
||||
.trimActions { display: flex; gap: 8px; flex-wrap: wrap; }
|
||||
.trimActions button { flex: 1; min-width: 120px; }
|
||||
|
||||
.taterLinkOverlay {
|
||||
position: fixed; inset: 0; padding: 22px;
|
||||
display: flex; align-items: center; justify-content: center;
|
||||
background: rgba(4,5,10,0.62); backdrop-filter: blur(12px);
|
||||
opacity: 0; visibility: hidden; pointer-events: none;
|
||||
transition: opacity 0.18s ease, visibility 0.18s ease;
|
||||
z-index: 12000;
|
||||
}
|
||||
.taterLinkOverlay.open { opacity: 1; visibility: visible; pointer-events: auto; }
|
||||
.taterLinkDialog {
|
||||
width: min(560px, calc(100vw - 36px));
|
||||
display: grid; gap: 18px; padding: 22px; border-radius: 24px;
|
||||
border: 1px solid rgba(255,138,42,0.28);
|
||||
background:
|
||||
radial-gradient(circle at top right, rgba(255,138,42,0.16), transparent 46%),
|
||||
linear-gradient(180deg, rgba(17,20,28,0.96), rgba(8,10,16,0.98));
|
||||
box-shadow: 0 30px 90px rgba(0,0,0,0.64);
|
||||
}
|
||||
.taterLinkCodePanel {
|
||||
display: grid; gap: 10px; text-align: center; padding: 24px;
|
||||
border-radius: 18px; border: 1px solid rgba(255,138,42,0.28);
|
||||
background: rgba(255,138,42,0.09);
|
||||
}
|
||||
.taterLinkCode {
|
||||
color: var(--orange2); font: 800 clamp(30px, 8vw, 46px)/1 ui-monospace, SFMono-Regular, Menlo, monospace;
|
||||
letter-spacing: 0.12em;
|
||||
}
|
||||
.taterLinkSuccess {
|
||||
display: grid; justify-items: center; gap: 12px; padding: 28px; text-align: center;
|
||||
}
|
||||
.taterLinkSuccessMark {
|
||||
display: grid; place-items: center; width: 68px; height: 68px; border-radius: 50%;
|
||||
background: rgba(57,212,160,0.15); border: 1px solid rgba(57,212,160,0.42);
|
||||
color: #6ee0af; font-size: 24px; font-weight: 900;
|
||||
}
|
||||
.taterLinkActions { display: flex; flex-wrap: wrap; gap: 10px; align-items: center; }
|
||||
|
||||
.pill.trimBadge {
|
||||
color: #89d4ff;
|
||||
border-color: rgba(137,212,255,0.25);
|
||||
@@ -1288,7 +1325,7 @@
|
||||
<div>
|
||||
<div class="studioKicker">False-Positive Loop</div>
|
||||
<h3>Auto Training</h3>
|
||||
<p>Transcribe real wake triggers, turn confirmed phrase-misses into hard negatives, retrain on your schedule, and ask Tater to refresh connected satellites.</p>
|
||||
<p>Sort false wakes into negatives, recover spoken close misses as positives, clean up confirmed wakes, retrain on your schedule, and refresh connected satellites.</p>
|
||||
<div class="studioSteps" aria-label="Auto Training steps">
|
||||
<span class="studioStepChip"><b>1</b> Transcribe wakes</span>
|
||||
<span class="studioStepChip"><b>2</b> Collect negatives</span>
|
||||
@@ -1305,13 +1342,21 @@
|
||||
<span class="studioStepBadge">1</span>
|
||||
<div>
|
||||
<h3>Review Rules</h3>
|
||||
<p>Only wake-trigger clips are reviewed automatically. Close misses, empty STT results, and clips containing the wake phrase remain in the inbox.</p>
|
||||
<p>False-wake sorting is always conservative. Cleanup and close-miss promotion remain optional, and empty STT results stay in the inbox.</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<label class="checkField">
|
||||
<input id="autoEnabled" type="checkbox" />
|
||||
<span><strong>Enable Auto Training</strong>New wake triggers will be queued for local Faster Whisper transcription.</span>
|
||||
<span><strong>Enable Auto Training</strong>Eligible wake triggers will be queued for transcription with the selected local STT engine.</span>
|
||||
</label>
|
||||
<label class="checkField">
|
||||
<input id="autoDeleteConfirmedWakes" type="checkbox" />
|
||||
<span><strong>Delete confirmed good wakes</strong>When a normal wake trigger contains the configured phrase, remove it from Captured Audio instead of keeping it for manual review.</span>
|
||||
</label>
|
||||
<label class="checkField">
|
||||
<input id="autoPromoteCloseMisses" type="checkbox" />
|
||||
<span><strong>Promote confirmed close misses</strong>Transcribe close misses that passed VAD and move them to Positive samples only when STT finds the configured phrase.</span>
|
||||
</label>
|
||||
<div class="autoGrid">
|
||||
<label class="field">
|
||||
@@ -1323,26 +1368,10 @@
|
||||
<input id="autoLanguage" type="text" value="en" placeholder="en" />
|
||||
</label>
|
||||
<label class="field wide">
|
||||
<strong>Faster Whisper model</strong>
|
||||
<input id="autoSttModel" type="text" value="small.en" placeholder="small.en" />
|
||||
</label>
|
||||
<label class="field">
|
||||
<strong>STT device</strong>
|
||||
<select id="autoSttDevice">
|
||||
<option value="auto" selected>Auto (prefer CUDA)</option>
|
||||
<option value="cuda">CUDA</option>
|
||||
<option value="cpu">CPU</option>
|
||||
</select>
|
||||
</label>
|
||||
<label class="field">
|
||||
<strong>Compute type</strong>
|
||||
<select id="autoSttComputeType">
|
||||
<option value="auto" selected>Auto (float16 CUDA / int8 CPU)</option>
|
||||
<option value="float16">float16</option>
|
||||
<option value="int8_float16">int8_float16</option>
|
||||
<option value="int8">int8</option>
|
||||
<option value="float32">float32</option>
|
||||
<option value="default">CTranslate2 default</option>
|
||||
<strong>STT engine</strong>
|
||||
<select id="autoSttEngine">
|
||||
<option value="faster_whisper" selected>Faster Whisper (recommended)</option>
|
||||
<option value="parakeet_onnx">Parakeet ONNX</option>
|
||||
</select>
|
||||
</label>
|
||||
<label class="field">
|
||||
@@ -1391,8 +1420,8 @@
|
||||
<div class="studioPanelTitle">
|
||||
<span class="studioStepBadge">3</span>
|
||||
<div>
|
||||
<h3>Publish + Satellite Refresh</h3>
|
||||
<p>The trainer publishes a LAN-reachable model URL, then asks Tater to re-push live settings so the firmware downloads the updated model at the same URL.</p>
|
||||
<h3>Publish to Tater</h3>
|
||||
<p>The trainer publishes a LAN-reachable model URL, then tells Tater to make the newly trained wake word active on every satellite.</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
@@ -1406,22 +1435,15 @@
|
||||
<strong>Tater URL</strong>
|
||||
<input id="autoTaterUrl" type="text" value="http://127.0.0.1:8501" />
|
||||
</label>
|
||||
<label class="field">
|
||||
<strong>Satellite selector (optional)</strong>
|
||||
<input id="autoTaterSelector" type="text" placeholder="Blank refreshes all connected sats" />
|
||||
</label>
|
||||
<label class="field">
|
||||
<strong>Tater API token (if enabled)</strong>
|
||||
<input id="autoTaterToken" type="password" placeholder="Not configured" autocomplete="off" />
|
||||
</label>
|
||||
</div>
|
||||
<div class="taterLinkActions">
|
||||
<span id="autoTaterLinkStatus" class="pill">Not linked</span>
|
||||
<button id="autoLinkTaterBtn" class="primary" type="button">Link Tater</button>
|
||||
<button id="autoUnlinkTaterBtn" class="danger" type="button" hidden>Unlink</button>
|
||||
</div>
|
||||
<label class="checkField">
|
||||
<input id="autoNotifySatellites" type="checkbox" checked />
|
||||
<span><strong>Refresh satellites after successful training</strong>Uses Tater's existing native satellite settings API.</span>
|
||||
</label>
|
||||
<label id="autoClearTokenRow" class="checkField" hidden>
|
||||
<input id="autoClearTaterToken" type="checkbox" />
|
||||
<span><strong>Clear the saved Tater token</strong>The token is otherwise preserved when the password field is blank.</span>
|
||||
<span><strong>Activate the new word after successful training</strong>Tater applies it globally and updates every connected satellite.</span>
|
||||
</label>
|
||||
</section>
|
||||
|
||||
@@ -1430,7 +1452,7 @@
|
||||
<button id="autoSaveBtn" class="primary" type="button">Save Auto Training</button>
|
||||
<button id="autoReviewNowBtn" type="button">Review inbox now</button>
|
||||
<button id="autoTrainNowBtn" type="button">Train now</button>
|
||||
<button id="autoNotifyNowBtn" type="button">Refresh satellites now</button>
|
||||
<button id="autoNotifyNowBtn" type="button">Publish current wake word now</button>
|
||||
</div>
|
||||
<div id="autoAudit" class="autoAudit muted">No automatic review has run yet.</div>
|
||||
</section>
|
||||
@@ -1673,6 +1695,22 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div id="taterLinkOverlay" class="taterLinkOverlay" aria-hidden="true">
|
||||
<div class="taterLinkDialog" role="dialog" aria-modal="true" aria-labelledby="taterLinkTitle">
|
||||
<div class="trimHeader">
|
||||
<div>
|
||||
<h3 id="taterLinkTitle" class="trimTitle">Link Tater</h3>
|
||||
<p id="taterLinkHint" class="trimHint">Enter the short-lived code shown in Tater Voice Settings.</p>
|
||||
</div>
|
||||
<button id="closeTaterLinkBtn" type="button">Close</button>
|
||||
</div>
|
||||
<div id="taterLinkBody">
|
||||
<div class="emptyState">Enter the secure pairing code from Tater.</div>
|
||||
</div>
|
||||
<div id="taterLinkModalStatus" class="muted">Waiting for the Tater code.</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<script>
|
||||
const $ = (id) => document.getElementById(id);
|
||||
|
||||
@@ -2030,28 +2068,33 @@
|
||||
const config = data.config || {};
|
||||
const state = data.state || {};
|
||||
const runtime = data.runtime || {};
|
||||
const trainerLink = data.trainer_link || {};
|
||||
uiState.autoTrain = data;
|
||||
|
||||
if (populateForm) {
|
||||
$("autoEnabled").checked = Boolean(config.enabled);
|
||||
$("autoWakePhrase").value = config.wake_phrase || uiState.session?.raw_phrase || "";
|
||||
$("autoLanguage").value = config.language || uiState.session?.language || "en";
|
||||
$("autoSttModel").value = config.stt_model || "small.en";
|
||||
$("autoSttDevice").value = config.stt_device || "auto";
|
||||
$("autoSttComputeType").value = config.stt_compute_type || "auto";
|
||||
$("autoSttEngine").value = config.stt_engine || "faster_whisper";
|
||||
$("autoMinimumChars").value = String(config.minimum_transcript_chars ?? 2);
|
||||
$("autoDeleteConfirmedWakes").checked = Boolean(config.delete_confirmed_wakes);
|
||||
$("autoPromoteCloseMisses").checked = Boolean(config.promote_close_misses);
|
||||
$("autoScheduleHours").value = String(config.schedule_hours ?? 24);
|
||||
$("autoMinimumNegatives").value = String(config.minimum_new_negatives ?? 3);
|
||||
$("autoAdvertisedUrl").value = config.advertised_base_url || "";
|
||||
$("autoTaterUrl").value = config.tater_url || "http://127.0.0.1:8501";
|
||||
$("autoTaterSelector").value = config.tater_selector || "";
|
||||
$("autoTaterToken").value = "";
|
||||
$("autoTaterToken").placeholder = config.tater_api_token_configured ? "Saved token (leave blank to keep)" : "Not configured";
|
||||
$("autoNotifySatellites").checked = config.notify_satellites !== false;
|
||||
$("autoClearTokenRow").hidden = !config.tater_api_token_configured;
|
||||
$("autoClearTaterToken").checked = false;
|
||||
}
|
||||
|
||||
const taterLinked = Boolean(trainerLink.linked);
|
||||
setPill(
|
||||
$("autoTaterLinkStatus"),
|
||||
taterLinked ? `Linked${trainerLink.tater_name ? ` to ${trainerLink.tater_name}` : ""}` : "Not linked",
|
||||
taterLinked ? "ok" : "warn"
|
||||
);
|
||||
$("autoLinkTaterBtn").textContent = taterLinked ? "Relink Tater" : "Link Tater";
|
||||
$("autoUnlinkTaterBtn").hidden = !taterLinked;
|
||||
|
||||
$("autoDetectedUrl").textContent = config.advertised_base_url
|
||||
? `Using configured URL: ${config.advertised_base_url}`
|
||||
: `Auto-detected URL: ${data.advertised_base_url || "unavailable"}`;
|
||||
@@ -2077,11 +2120,14 @@
|
||||
if (state.last_review_file) audit.push(state.last_review_file);
|
||||
if (state.last_review_transcript) audit.push(`STT: “${state.last_review_transcript}”`);
|
||||
if (state.last_review_error) audit.push(`Error: ${state.last_review_error}`);
|
||||
if (state.last_stt_device) audit.push(`STT runtime: ${state.last_stt_device} / ${state.last_stt_compute_type || "default"}`);
|
||||
if (state.last_stt_engine) {
|
||||
const runtimeLabel = [state.last_stt_device, state.last_stt_compute_type].filter(Boolean).join(" / ");
|
||||
audit.push(`STT engine: ${String(state.last_stt_engine).replaceAll("_", " ")}${runtimeLabel ? ` · ${runtimeLabel}` : ""}`);
|
||||
}
|
||||
if (state.last_notify_at) {
|
||||
audit.push(state.last_notify_error
|
||||
? `Satellite refresh failed: ${state.last_notify_error}`
|
||||
: `Satellite refresh: ${state.last_notify_count ?? "requested"} connected at ${formatTimestamp(state.last_notify_at)}`);
|
||||
? `Wake-word publish failed: ${state.last_notify_error}`
|
||||
: `Wake word published to ${state.last_notify_count ?? "all"} connected satellite(s) at ${formatTimestamp(state.last_notify_at)}`);
|
||||
}
|
||||
$("autoAudit").textContent = audit.join(" · ") || "No automatic review has run yet.";
|
||||
syncButtons();
|
||||
@@ -2098,20 +2144,16 @@
|
||||
enabled: $("autoEnabled").checked,
|
||||
wake_phrase: ($("autoWakePhrase").value || "").trim(),
|
||||
language: ($("autoLanguage").value || "en").trim(),
|
||||
stt_model: ($("autoSttModel").value || "").trim(),
|
||||
stt_device: $("autoSttDevice").value || "auto",
|
||||
stt_compute_type: $("autoSttComputeType").value || "auto",
|
||||
stt_engine: $("autoSttEngine").value || "faster_whisper",
|
||||
minimum_transcript_chars: Number($("autoMinimumChars").value || 2),
|
||||
delete_confirmed_wakes: $("autoDeleteConfirmedWakes").checked,
|
||||
promote_close_misses: $("autoPromoteCloseMisses").checked,
|
||||
schedule_hours: Number($("autoScheduleHours").value || 0),
|
||||
minimum_new_negatives: Number($("autoMinimumNegatives").value || 3),
|
||||
advertised_base_url: ($("autoAdvertisedUrl").value || "").trim(),
|
||||
tater_url: ($("autoTaterUrl").value || "").trim(),
|
||||
tater_selector: ($("autoTaterSelector").value || "").trim(),
|
||||
notify_satellites: $("autoNotifySatellites").checked,
|
||||
clear_tater_api_token: $("autoClearTaterToken").checked,
|
||||
};
|
||||
const token = ($("autoTaterToken").value || "").trim();
|
||||
if (token) payload.tater_api_token = token;
|
||||
return payload;
|
||||
}
|
||||
|
||||
@@ -2153,7 +2195,7 @@
|
||||
await refreshSession();
|
||||
pollTraining();
|
||||
} else {
|
||||
setPill($("autoStatus"), `Satellite refresh requested${data.count === null || data.count === undefined ? "" : ` for ${data.count}`}`, "ok");
|
||||
setPill($("autoStatus"), `Wake word published${data.count === null || data.count === undefined ? "" : ` to ${data.count} satellite(s)`}`, "ok");
|
||||
}
|
||||
return data;
|
||||
} finally {
|
||||
@@ -2162,6 +2204,98 @@
|
||||
}
|
||||
}
|
||||
|
||||
function closeTaterLinkModal() {
|
||||
$("taterLinkOverlay").classList.remove("open");
|
||||
$("taterLinkOverlay").setAttribute("aria-hidden", "true");
|
||||
}
|
||||
|
||||
function showTaterLinkSuccess(status) {
|
||||
const taterName = status?.tater_name ? ` to ${escapeHtml(status.tater_name)}` : "";
|
||||
$("taterLinkTitle").textContent = "Tater linked";
|
||||
$("taterLinkHint").textContent = "This trainer can now securely publish wake-word updates.";
|
||||
$("taterLinkBody").innerHTML = `
|
||||
<div class="taterLinkSuccess">
|
||||
<div class="taterLinkSuccessMark" aria-hidden="true">✓</div>
|
||||
<strong>Successfully linked${taterName}</strong>
|
||||
<span class="muted">The private link key is stored locally and is never shown again.</span>
|
||||
</div>
|
||||
`;
|
||||
$("taterLinkModalStatus").textContent = "You can close this popup.";
|
||||
}
|
||||
|
||||
async function openTaterLinkModal() {
|
||||
$("taterLinkTitle").textContent = "Link Tater";
|
||||
$("taterLinkHint").textContent = "Enter the short-lived code shown in Tater Voice Settings.";
|
||||
const taterUrl = ($("autoTaterUrl").value || "http://127.0.0.1:8501").trim();
|
||||
$("taterLinkBody").innerHTML = `
|
||||
<div class="stack">
|
||||
<label class="field">
|
||||
<strong>Tater address</strong>
|
||||
<input id="taterLinkUrl" type="text" value="${escapeAttr(taterUrl)}" placeholder="http://127.0.0.1:8501" />
|
||||
</label>
|
||||
<div class="taterLinkCodePanel">
|
||||
<span class="muted">Tater pairing code</span>
|
||||
<input id="taterLinkCode" class="taterLinkCode" type="text" inputmode="text" autocomplete="off"
|
||||
maxlength="9" placeholder="ABCD-EFGH" />
|
||||
<span class="muted">In Tater, open Voice Settings → Wake Word Trainer → Link Trainer.</span>
|
||||
</div>
|
||||
<button id="claimTaterLinkBtn" class="primary" type="button">Link Tater</button>
|
||||
</div>
|
||||
`;
|
||||
$("taterLinkModalStatus").textContent = "Waiting for the Tater code.";
|
||||
$("taterLinkOverlay").classList.add("open");
|
||||
$("taterLinkOverlay").setAttribute("aria-hidden", "false");
|
||||
const codeInput = $("taterLinkCode");
|
||||
const urlInput = $("taterLinkUrl");
|
||||
const submit = $("claimTaterLinkBtn");
|
||||
codeInput.addEventListener("input", () => {
|
||||
const raw = String(codeInput.value || "").toUpperCase().replace(/[^A-Z0-9]/g, "").slice(0, 8);
|
||||
codeInput.value = raw.length > 4 ? `${raw.slice(0, 4)}-${raw.slice(4)}` : raw;
|
||||
});
|
||||
submit.addEventListener("click", async () => {
|
||||
const pairingCode = String(codeInput.value || "").trim();
|
||||
const targetUrl = String(urlInput.value || "").trim();
|
||||
if (!pairingCode || !targetUrl) {
|
||||
$("taterLinkModalStatus").textContent = "Tater address and pairing code are required.";
|
||||
return;
|
||||
}
|
||||
submit.disabled = true;
|
||||
codeInput.disabled = true;
|
||||
urlInput.disabled = true;
|
||||
$("taterLinkModalStatus").textContent = "Linking securely with Tater...";
|
||||
try {
|
||||
const result = await api("/api/tater_link/claim", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ tater_url: targetUrl, pairing_code: pairingCode }),
|
||||
});
|
||||
$("autoTaterUrl").value = targetUrl;
|
||||
showTaterLinkSuccess(result);
|
||||
await refreshAutoTrain(false);
|
||||
} catch (error) {
|
||||
submit.disabled = false;
|
||||
codeInput.disabled = false;
|
||||
urlInput.disabled = false;
|
||||
$("taterLinkModalStatus").textContent = `Link failed: ${error.message}`;
|
||||
}
|
||||
});
|
||||
window.setTimeout(() => codeInput.focus(), 50);
|
||||
}
|
||||
|
||||
async function unlinkTater() {
|
||||
if (!window.confirm("Unlink this trainer from Tater? Wake-word publishing will stop until it is linked again.")) return;
|
||||
uiState.autoBusy = true;
|
||||
syncButtons();
|
||||
try {
|
||||
await api("/api/tater_link/unlink", { method: "POST" });
|
||||
await refreshAutoTrain(false);
|
||||
setPill($("autoTaterLinkStatus"), "Not linked", "warn");
|
||||
} finally {
|
||||
uiState.autoBusy = false;
|
||||
syncButtons();
|
||||
}
|
||||
}
|
||||
|
||||
function captureBadge(item) {
|
||||
if (item.blocked_by_vad) return { label: "Blocked by VAD", cls: "warn" };
|
||||
const eventType = String(item?.event_type || "").toLowerCase();
|
||||
@@ -2207,6 +2341,7 @@
|
||||
if (item.average_probability !== null && item.average_probability !== undefined) meta.push(`<span class="pill">avg ${escapeHtml(item.average_probability)}</span>`);
|
||||
if (item.detection_profile) meta.push(`<span class="pill">profile ${escapeHtml(formatDetectionProfile(item.detection_profile))}</span>`);
|
||||
if (item.auto_review_status) meta.push(`<span class="pill ${item.auto_review_status === "error" ? "err" : "warn"}">auto ${escapeHtml(String(item.auto_review_status).replaceAll("_", " "))}</span>`);
|
||||
if (item.auto_review_match_method === "guided_close_match") meta.push(`<span class="pill ok">guided wake confirmation</span>`);
|
||||
if (item.peak_probability_cutoff !== null && item.peak_probability_cutoff !== undefined) meta.push(`<span class="pill">peak cutoff ${escapeHtml(item.peak_probability_cutoff)}</span>`);
|
||||
if (item.probability_cutoff !== null && item.probability_cutoff !== undefined) meta.push(`<span class="pill">avg cutoff ${escapeHtml(item.probability_cutoff)}</span>`);
|
||||
if (item.active_window_count !== null && item.active_window_count !== undefined && item.min_active_windows !== null && item.min_active_windows !== undefined) {
|
||||
@@ -2235,6 +2370,8 @@
|
||||
</div>
|
||||
<div class="fileMeta">${meta.join("") || `<span class="muted">No metadata attached</span>`}</div>
|
||||
${item.transcript ? `<div class="autoAudit"><strong>STT transcript</strong><br>${escapeHtml(item.transcript)}</div>` : ""}
|
||||
${item.auto_review_guided_transcript ? `<div class="autoAudit"><strong>Guided wake check</strong><br>${escapeHtml(item.auto_review_guided_transcript)}</div>` : ""}
|
||||
${item.auto_review_reason ? `<div class="muted">${escapeHtml(item.auto_review_reason)}</div>` : ""}
|
||||
${item.auto_review_error ? `<div class="muted">Auto review error: ${escapeHtml(item.auto_review_error)}</div>` : ""}
|
||||
<audio class="audioPlayer" controls preload="none" src="${escapeHtml(item.audio_url || `/api/audio/captured/${encodeURIComponent(item.saved_as)}`)}"></audio>
|
||||
<div class="muted">Stored as ${escapeHtml(item.saved_as)} · ${escapeHtml(formatSummary)}</div>
|
||||
@@ -2280,6 +2417,7 @@
|
||||
if (when) subtitleParts.push(`Saved ${when}`);
|
||||
if (item.message) subtitleParts.push(item.message);
|
||||
if (item.auto_negative) subtitleParts.push("Auto-reviewed false positive");
|
||||
if (item.auto_positive) subtitleParts.push("Auto-promoted close miss");
|
||||
let revertBtn = '';
|
||||
if (item.trimmed) {
|
||||
revertBtn = `<button type="button" data-sample-revert="${escapeAttr(item.saved_as)}" data-bucket="${escapeAttr(bucket)}">Revert</button>`;
|
||||
@@ -2582,9 +2720,14 @@
|
||||
if (refreshWakeWordsBtn) {
|
||||
refreshWakeWordsBtn.disabled = uiState.firmwareBusy;
|
||||
}
|
||||
for (const id of ["autoSaveBtn", "autoReviewNowBtn", "autoTrainNowBtn", "autoNotifyNowBtn"]) {
|
||||
for (const id of ["autoSaveBtn", "autoReviewNowBtn", "autoTrainNowBtn", "autoNotifyNowBtn", "autoLinkTaterBtn", "autoUnlinkTaterBtn"]) {
|
||||
const button = $(id);
|
||||
if (button) button.disabled = uiState.autoBusy || (id === "autoTrainNowBtn" && Boolean(training.running));
|
||||
if (button) {
|
||||
button.disabled =
|
||||
uiState.autoBusy ||
|
||||
(id === "autoTrainNowBtn" && Boolean(training.running)) ||
|
||||
(id === "autoNotifyNowBtn" && !Boolean(uiState.autoTrain?.trainer_link?.linked));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2927,10 +3070,26 @@
|
||||
try {
|
||||
await runAutoTrainAction("notify_now");
|
||||
} catch (error) {
|
||||
setPill($("autoStatus"), "Refresh failed", "err");
|
||||
alert("Satellite refresh failed: " + error.message);
|
||||
setPill($("autoStatus"), "Publish failed", "err");
|
||||
alert("Wake-word publish failed: " + error.message);
|
||||
}
|
||||
});
|
||||
$("autoLinkTaterBtn").addEventListener("click", () => {
|
||||
openTaterLinkModal().catch((error) => {
|
||||
setPill($("autoTaterLinkStatus"), "Link failed", "err");
|
||||
alert("Tater link failed: " + error.message);
|
||||
});
|
||||
});
|
||||
$("autoUnlinkTaterBtn").addEventListener("click", () => {
|
||||
unlinkTater().catch((error) => {
|
||||
setPill($("autoTaterLinkStatus"), "Unlink failed", "err");
|
||||
alert("Tater unlink failed: " + error.message);
|
||||
});
|
||||
});
|
||||
$("closeTaterLinkBtn").addEventListener("click", closeTaterLinkModal);
|
||||
$("taterLinkOverlay").addEventListener("click", (event) => {
|
||||
if (event.target === $("taterLinkOverlay")) closeTaterLinkModal();
|
||||
});
|
||||
|
||||
$("openConsoleBtn").addEventListener("click", () => {
|
||||
setConsoleLogAutoScroll($("trainLog"), (uiState.training?.log_lines || []).join("\n") || "(no training started)");
|
||||
@@ -2949,6 +3108,7 @@
|
||||
|
||||
document.addEventListener("keydown", (event) => {
|
||||
if (event.key === "Escape") {
|
||||
closeTaterLinkModal();
|
||||
closeConsole();
|
||||
}
|
||||
});
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import io
|
||||
import json
|
||||
import queue
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
@@ -22,7 +23,18 @@ def silent_wav_bytes(duration_s: float = 0.25) -> bytes:
|
||||
|
||||
|
||||
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 = (
|
||||
@@ -31,12 +43,14 @@ class AutoTrainTests(unittest.TestCase):
|
||||
trainer.PERSONAL_DIR,
|
||||
trainer.AUTO_TRAIN_CONFIG_FILE,
|
||||
trainer.AUTO_TRAIN_STATE_FILE,
|
||||
trainer.AUTO_TRAIN_MODEL_DIR,
|
||||
)
|
||||
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"
|
||||
trainer.AUTO_TRAIN_MODEL_DIR = root / "auto_train_models"
|
||||
for directory in (trainer.CAPTURED_DIR, trainer.NEGATIVE_DIR, trainer.PERSONAL_DIR):
|
||||
directory.mkdir(parents=True)
|
||||
|
||||
@@ -50,8 +64,6 @@ class AutoTrainTests(unittest.TestCase):
|
||||
"wake_phrase": "hey tater",
|
||||
"language": "en",
|
||||
"tater_url": "http://127.0.0.1:8501",
|
||||
"stt_device": "auto",
|
||||
"stt_compute_type": "auto",
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -65,14 +77,22 @@ class AutoTrainTests(unittest.TestCase):
|
||||
trainer.PERSONAL_DIR,
|
||||
trainer.AUTO_TRAIN_CONFIG_FILE,
|
||||
trainer.AUTO_TRAIN_STATE_FILE,
|
||||
trainer.AUTO_TRAIN_MODEL_DIR,
|
||||
) = 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") -> Path:
|
||||
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(
|
||||
@@ -80,7 +100,8 @@ class AutoTrainTests(unittest.TestCase):
|
||||
{
|
||||
"original_name": name,
|
||||
"wake_word": wake_word,
|
||||
"event_type": "wake_detected",
|
||||
"event_type": event_type,
|
||||
"blocked_by_vad": blocked_by_vad,
|
||||
"review_status": "pending",
|
||||
},
|
||||
)
|
||||
@@ -90,9 +111,199 @@ class AutoTrainTests(unittest.TestCase):
|
||||
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_similarity_recognizes_real_short_clip_mishearings(self):
|
||||
for transcript in ("Hey, haters.", "Hate hater.", "Hey Ganger.", "Hey, gator."):
|
||||
with self.subTest(transcript=transcript):
|
||||
self.assertGreaterEqual(
|
||||
trainer._wake_phrase_similarity(transcript, "hey tater"),
|
||||
trainer.WAKE_PHRASE_GUIDANCE_MIN_SIMILARITY,
|
||||
)
|
||||
|
||||
for transcript in ("turn on the lights", "what is the weather", "play some music"):
|
||||
with self.subTest(transcript=transcript):
|
||||
self.assertLess(
|
||||
trainer._wake_phrase_similarity(transcript, "hey tater"),
|
||||
trainer.WAKE_PHRASE_GUIDANCE_MIN_SIMILARITY,
|
||||
)
|
||||
|
||||
def test_stt_engine_selection_uses_managed_models(self):
|
||||
config = trainer._normalize_auto_train_config(
|
||||
{
|
||||
"stt_engine": "parakeet-onnx",
|
||||
"stt_model": "user/should-not-be-used",
|
||||
"stt_device": "cpu",
|
||||
"stt_compute_type": "float32",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(config["stt_engine"], trainer.STT_ENGINE_PARAKEET_ONNX)
|
||||
self.assertNotIn("stt_model", config)
|
||||
self.assertNotIn("stt_device", config)
|
||||
self.assertNotIn("stt_compute_type", config)
|
||||
self.assertEqual(
|
||||
trainer._managed_stt_model(config["stt_engine"], "en"),
|
||||
trainer.DEFAULT_PARAKEET_ONNX_MODEL,
|
||||
)
|
||||
self.assertEqual(
|
||||
trainer._managed_stt_model(trainer.STT_ENGINE_FASTER_WHISPER, "de"),
|
||||
trainer.DEFAULT_FASTER_WHISPER_MULTILINGUAL_MODEL,
|
||||
)
|
||||
|
||||
def test_stt_router_supports_both_nvidia_engines(self):
|
||||
audio_path = Path("wake.wav")
|
||||
with (
|
||||
patch.object(trainer, "_transcribe_capture_with_faster_whisper", return_value="faster") as faster,
|
||||
patch.object(trainer, "_transcribe_capture_with_parakeet", return_value="parakeet") as parakeet,
|
||||
):
|
||||
self.assertEqual(
|
||||
trainer._transcribe_capture(
|
||||
audio_path,
|
||||
engine=trainer.STT_ENGINE_FASTER_WHISPER,
|
||||
language="en",
|
||||
),
|
||||
"faster",
|
||||
)
|
||||
self.assertEqual(
|
||||
trainer._transcribe_capture(
|
||||
audio_path,
|
||||
engine=trainer.STT_ENGINE_PARAKEET_ONNX,
|
||||
language="en",
|
||||
),
|
||||
"parakeet",
|
||||
)
|
||||
|
||||
faster.assert_called_once()
|
||||
parakeet.assert_called_once()
|
||||
|
||||
def test_guided_faster_whisper_uses_dynamic_wake_phrase(self):
|
||||
fake_model = SimpleNamespace(
|
||||
transcribe=Mock(
|
||||
return_value=(
|
||||
iter([SimpleNamespace(text=" hello "), SimpleNamespace(text="potato ")]),
|
||||
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_guided(
|
||||
Path("wake.wav"),
|
||||
model="small.en",
|
||||
language="en",
|
||||
wake_phrase="Hello_Potato",
|
||||
)
|
||||
|
||||
self.assertEqual(transcript, "hello potato")
|
||||
_, kwargs = fake_model.transcribe.call_args
|
||||
self.assertEqual(kwargs["hotwords"], "hello potato")
|
||||
self.assertIn("hello potato", kwargs["initial_prompt"])
|
||||
self.assertEqual(kwargs["beam_size"], 5)
|
||||
self.assertEqual(kwargs["best_of"], 5)
|
||||
self.assertEqual(kwargs["temperature"], 0.0)
|
||||
self.assertFalse(kwargs["condition_on_previous_text"])
|
||||
|
||||
def test_parakeet_loader_prefers_cuda_then_cpu(self):
|
||||
fake_model = object()
|
||||
fake_onnx_asr = SimpleNamespace(load_model=Mock(return_value=fake_model))
|
||||
fake_huggingface_hub = SimpleNamespace(
|
||||
snapshot_download=Mock(return_value=str(trainer.AUTO_TRAIN_MODEL_DIR))
|
||||
)
|
||||
with (
|
||||
patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"onnx_asr": fake_onnx_asr,
|
||||
"huggingface_hub": fake_huggingface_hub,
|
||||
},
|
||||
),
|
||||
patch.object(
|
||||
trainer,
|
||||
"_parakeet_onnx_providers",
|
||||
return_value=["CUDAExecutionProvider", "CPUExecutionProvider"],
|
||||
),
|
||||
):
|
||||
with trainer.PARAKEET_ONNX_MODEL_LOCK:
|
||||
trainer.PARAKEET_ONNX_MODEL_CACHE.clear()
|
||||
loaded = trainer._load_parakeet_onnx_model()
|
||||
|
||||
self.assertIs(loaded, fake_model)
|
||||
fake_huggingface_hub.snapshot_download.assert_called_once_with(
|
||||
repo_id=trainer.DEFAULT_PARAKEET_ONNX_REPO,
|
||||
local_dir=str(trainer.AUTO_TRAIN_MODEL_DIR),
|
||||
allow_patterns=[
|
||||
"config.json",
|
||||
"vocab.txt",
|
||||
"encoder-model.int8.onnx",
|
||||
"encoder-model.int8.onnx.data",
|
||||
"decoder_joint-model.int8.onnx",
|
||||
"decoder_joint-model.int8.onnx.data",
|
||||
],
|
||||
)
|
||||
fake_onnx_asr.load_model.assert_called_once_with(
|
||||
trainer.DEFAULT_PARAKEET_ONNX_MODEL,
|
||||
str(trainer.AUTO_TRAIN_MODEL_DIR),
|
||||
quantization="int8",
|
||||
providers=["CUDAExecutionProvider", "CPUExecutionProvider"],
|
||||
)
|
||||
|
||||
def test_parakeet_loader_reuses_complete_snapshot_offline(self):
|
||||
fake_model = object()
|
||||
fake_onnx_asr = SimpleNamespace(load_model=Mock(return_value=fake_model))
|
||||
fake_huggingface_hub = SimpleNamespace(snapshot_download=Mock())
|
||||
trainer.AUTO_TRAIN_MODEL_DIR.mkdir(parents=True, exist_ok=True)
|
||||
for filename in (
|
||||
"config.json",
|
||||
"vocab.txt",
|
||||
"encoder-model.int8.onnx",
|
||||
"decoder_joint-model.int8.onnx",
|
||||
):
|
||||
(trainer.AUTO_TRAIN_MODEL_DIR / filename).touch()
|
||||
with (
|
||||
patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"onnx_asr": fake_onnx_asr,
|
||||
"huggingface_hub": fake_huggingface_hub,
|
||||
},
|
||||
),
|
||||
patch.object(
|
||||
trainer,
|
||||
"_parakeet_onnx_providers",
|
||||
return_value=["CUDAExecutionProvider", "CPUExecutionProvider"],
|
||||
),
|
||||
):
|
||||
with trainer.PARAKEET_ONNX_MODEL_LOCK:
|
||||
trainer.PARAKEET_ONNX_MODEL_CACHE.clear()
|
||||
loaded = trainer._load_parakeet_onnx_model()
|
||||
|
||||
self.assertIs(loaded, fake_model)
|
||||
fake_huggingface_hub.snapshot_download.assert_not_called()
|
||||
fake_onnx_asr.load_model.assert_called_once_with(
|
||||
trainer.DEFAULT_PARAKEET_ONNX_MODEL,
|
||||
str(trainer.AUTO_TRAIN_MODEL_DIR),
|
||||
quantization="int8",
|
||||
providers=["CUDAExecutionProvider", "CPUExecutionProvider"],
|
||||
)
|
||||
|
||||
def test_ui_exposes_engine_selector_without_manual_runtime_fields(self):
|
||||
source = (Path(__file__).resolve().parents[1] / "static" / "index.html").read_text(
|
||||
encoding="utf-8"
|
||||
)
|
||||
self.assertIn('id="autoSttEngine"', source)
|
||||
self.assertNotIn('id="autoSttModel"', source)
|
||||
self.assertNotIn('id="autoSttDevice"', source)
|
||||
self.assertNotIn('id="autoSttComputeType"', source)
|
||||
self.assertIn("Guided wake check", source)
|
||||
|
||||
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"):
|
||||
with patch.object(trainer, "_transcribe_capture", return_value="turn on the kitchen lights"):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertFalse((trainer.CAPTURED_DIR / "wake.wav").exists())
|
||||
@@ -102,11 +313,13 @@ class AutoTrainTests(unittest.TestCase):
|
||||
self.assertTrue(metadata["auto_negative"])
|
||||
self.assertEqual(metadata["review_status"], "auto_approved_negative")
|
||||
self.assertEqual(metadata["transcript"], "turn on the kitchen lights")
|
||||
self.assertEqual(metadata["auto_review_stt_engine"], "faster_whisper")
|
||||
self.assertEqual(metadata["auto_review_stt_model"], "small.en")
|
||||
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"):
|
||||
with patch.object(trainer, "_transcribe_capture", return_value="hey tater turn on the lights"):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
self.assertTrue(audio_path.exists())
|
||||
@@ -115,9 +328,177 @@ class AutoTrainTests(unittest.TestCase):
|
||||
self.assertEqual(metadata["auto_review_status"], "wake_phrase_detected")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["pending_negative_count"], 0)
|
||||
|
||||
def test_close_transcript_uses_guided_faster_whisper_confirmation(self):
|
||||
audio_path = self.add_capture()
|
||||
with (
|
||||
patch.object(trainer, "_transcribe_capture", return_value="Hey, haters."),
|
||||
patch.object(
|
||||
trainer,
|
||||
"_transcribe_capture_with_faster_whisper_guided",
|
||||
return_value="Hey Tater",
|
||||
) as guided,
|
||||
):
|
||||
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(metadata["transcript"], "Hey, haters.")
|
||||
self.assertEqual(metadata["auto_review_guided_transcript"], "Hey Tater")
|
||||
self.assertEqual(metadata["auto_review_match_method"], "guided_close_match")
|
||||
self.assertGreaterEqual(
|
||||
metadata["auto_review_phrase_similarity"],
|
||||
trainer.WAKE_PHRASE_GUIDANCE_MIN_SIMILARITY,
|
||||
)
|
||||
guided.assert_called_once()
|
||||
guided_args, guided_kwargs = guided.call_args
|
||||
self.assertEqual(guided_args[0].resolve(), audio_path.resolve())
|
||||
self.assertEqual(
|
||||
guided_kwargs,
|
||||
{
|
||||
"model": "small.en",
|
||||
"language": "en",
|
||||
"wake_phrase": "hey tater",
|
||||
},
|
||||
)
|
||||
|
||||
def test_unconfirmed_close_transcript_stays_for_manual_review(self):
|
||||
audio_path = self.add_capture()
|
||||
with (
|
||||
patch.object(trainer, "_transcribe_capture", return_value="Hate hater."),
|
||||
patch.object(
|
||||
trainer,
|
||||
"_transcribe_capture_with_faster_whisper_guided",
|
||||
return_value="Hate hater.",
|
||||
),
|
||||
):
|
||||
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_ambiguous")
|
||||
self.assertEqual(metadata["transcript"], "Hate hater.")
|
||||
self.assertEqual(metadata["auto_review_guided_transcript"], "Hate hater.")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["pending_negative_count"], 0)
|
||||
|
||||
self.assertEqual(trainer._queue_pending_auto_reviews(), 0)
|
||||
self.assertEqual(trainer._queue_pending_auto_reviews(force=True), 1)
|
||||
|
||||
def test_close_parakeet_transcript_stays_for_manual_review(self):
|
||||
audio_path = self.add_capture()
|
||||
trainer.AUTO_TRAIN_CONFIG["stt_engine"] = trainer.STT_ENGINE_PARAKEET_ONNX
|
||||
with (
|
||||
patch.object(trainer, "_transcribe_capture", return_value="Hey Ganger."),
|
||||
patch.object(trainer, "_transcribe_capture_with_faster_whisper_guided") as guided,
|
||||
):
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
guided.assert_not_called()
|
||||
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_ambiguous")
|
||||
self.assertEqual(metadata["auto_review_stt_engine"], "parakeet_onnx")
|
||||
|
||||
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",
|
||||
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") 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") 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", 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",
|
||||
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") 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:
|
||||
with patch.object(trainer, "_transcribe_capture") as transcribe:
|
||||
trainer._auto_review_capture("wake.wav")
|
||||
|
||||
transcribe.assert_not_called()
|
||||
@@ -136,13 +517,12 @@ class AutoTrainTests(unittest.TestCase):
|
||||
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):
|
||||
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_selector": "kitchen-sat",
|
||||
"tater_api_token": "secret-token",
|
||||
"tater_link_token": "secret-token",
|
||||
}
|
||||
)
|
||||
|
||||
@@ -154,17 +534,96 @@ class AutoTrainTests(unittest.TestCase):
|
||||
return False
|
||||
|
||||
def read(self):
|
||||
return b'{"push":{"count":2}}'
|
||||
return b'{"push":{"count":4}}'
|
||||
|
||||
with patch.object(trainer, "urlopen", return_value=Response()) as open_url:
|
||||
result = trainer._notify_tater_satellites()
|
||||
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"], 2)
|
||||
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/settings")
|
||||
self.assertEqual(request.get_header("X-tater-token"), "secret-token")
|
||||
self.assertEqual(json.loads(request.data), {"selector": "kitchen-sat", "settings": {}})
|
||||
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(
|
||||
@@ -233,6 +692,39 @@ class AutoTrainTests(unittest.TestCase):
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_stt_device"], "cuda")
|
||||
self.assertEqual(trainer.AUTO_TRAIN_STATE["last_stt_compute_type"], "float16")
|
||||
|
||||
def test_train_status_reads_and_increments_training_log_tail(self):
|
||||
log_path = Path(self.tempdir.name) / "training.log"
|
||||
log_path.write_text("first\nsecond\nthird\n", encoding="utf-8")
|
||||
with trainer.STATE_LOCK:
|
||||
original_training = dict(trainer.STATE["training"])
|
||||
trainer.STATE["training"].update(
|
||||
{
|
||||
"log_path": str(log_path),
|
||||
"last_sent_tail": [],
|
||||
"last_log_size": 0,
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(trainer, "TRAIN_LOG_TAIL_LINES", 2),
|
||||
patch.object(trainer, "TRAIN_LOG_MAX_BYTES", 1024),
|
||||
):
|
||||
first_status = trainer.train_status()
|
||||
self.assertEqual(first_status["training"]["log_lines"], ["second", "third"])
|
||||
self.assertEqual(first_status["training"]["log_text"], "second\nthird")
|
||||
|
||||
with log_path.open("a", encoding="utf-8") as log_file:
|
||||
log_file.write("fourth\n")
|
||||
|
||||
next_status = trainer.train_status()
|
||||
self.assertEqual(next_status["training"]["log_lines"], ["third", "fourth"])
|
||||
self.assertEqual(next_status["training"]["log_text"], "fourth")
|
||||
finally:
|
||||
with trainer.STATE_LOCK:
|
||||
trainer.STATE["training"].clear()
|
||||
trainer.STATE["training"].update(original_training)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
73
tests/test_run_sh.py
Normal file
73
tests/test_run_sh.py
Normal file
@@ -0,0 +1,73 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
RUN_SH = REPO_ROOT / "run.sh"
|
||||
|
||||
|
||||
def _cuda_path_probe() -> str:
|
||||
source = RUN_SH.read_text(encoding="utf-8")
|
||||
match = re.search(
|
||||
r'WHISPER_CUDA_LIBRARY_PATH="\$\("\$\{PY\}" - <<\'PY\'\n(?P<probe>.*?)\nPY\n\)"',
|
||||
source,
|
||||
flags=re.DOTALL,
|
||||
)
|
||||
if match is None:
|
||||
raise AssertionError("Could not locate the CUDA library path probe in run.sh")
|
||||
return match.group("probe")
|
||||
|
||||
|
||||
class RunShCudaLibraryPathTests(unittest.TestCase):
|
||||
def _run_probe(self, python_path: Path) -> subprocess.CompletedProcess[str]:
|
||||
env = dict(os.environ)
|
||||
env["PYTHONPATH"] = str(python_path)
|
||||
return subprocess.run(
|
||||
[sys.executable, "-S", "-"],
|
||||
input=_cuda_path_probe(),
|
||||
text=True,
|
||||
capture_output=True,
|
||||
check=False,
|
||||
env=env,
|
||||
)
|
||||
|
||||
def test_namespace_cuda_packages_do_not_require_module_file(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
root = Path(temp_dir)
|
||||
cublas_lib = root / "nvidia" / "cublas" / "lib"
|
||||
cudnn_lib = root / "nvidia" / "cudnn" / "lib"
|
||||
cublas_lib.mkdir(parents=True)
|
||||
cudnn_lib.mkdir(parents=True)
|
||||
|
||||
result = self._run_probe(root)
|
||||
|
||||
self.assertEqual(result.returncode, 0, result.stderr)
|
||||
self.assertEqual(
|
||||
result.stdout.strip().split(":"),
|
||||
[str(cublas_lib.resolve()), str(cudnn_lib.resolve())],
|
||||
)
|
||||
|
||||
def test_missing_cuda_packages_return_an_empty_path(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
result = self._run_probe(Path(temp_dir))
|
||||
|
||||
self.assertEqual(result.returncode, 0, result.stderr)
|
||||
self.assertEqual(result.stdout.strip(), "")
|
||||
|
||||
def test_parakeet_uses_cuda_onnxruntime_package(self) -> None:
|
||||
source = RUN_SH.read_text(encoding="utf-8")
|
||||
|
||||
self.assertIn('"onnx-asr[hub]>=0.12.0"', source)
|
||||
self.assertIn('"onnxruntime-gpu[cuda,cudnn]<1.27"', source)
|
||||
self.assertIn('"CUDAExecutionProvider"', source)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user