Release NVIDIA WakeWord Trainer v23

This commit is contained in:
MasterPhooey
2026-08-03 06:45:32 -05:00
parent 2a88090b85
commit 1f16f6f916
4 changed files with 68 additions and 22 deletions

View File

@@ -1,7 +1,11 @@
from __future__ import annotations
import io
import signal
import tempfile
import threading
import unittest
from pathlib import Path
from unittest.mock import patch
import trainer_server as trainer
@@ -26,6 +30,19 @@ class _FakeTrainingProcess:
self.returncode = -signal.SIGKILL
class _CompletedTrainingProcess:
def __init__(self):
self.pid = 6543
self.returncode = 0
self.stdout = io.StringIO("worker started\n")
def poll(self):
return self.returncode
def wait(self, timeout=None):
return self.returncode
class SessionStopTests(unittest.TestCase):
def tearDown(self):
trainer.TRAINING_STOP_EVENT.clear()
@@ -49,6 +66,51 @@ class SessionStopTests(unittest.TestCase):
trainer.TRAINING_PROCESS = original_process
trainer.TRAINING_THREAD = original_thread
def test_reserved_running_state_starts_the_background_worker(self):
original_process = trainer.TRAINING_PROCESS
original_thread = trainer.TRAINING_THREAD
original_raw_phrase = trainer.STATE.get("raw_phrase")
original_training = dict(trainer.STATE["training"])
try:
with tempfile.TemporaryDirectory() as directory:
data_dir = Path(directory)
process = _CompletedTrainingProcess()
trainer.TRAINING_PROCESS = None
trainer.TRAINING_THREAD = threading.current_thread()
with trainer.STATE_LOCK:
trainer.STATE["raw_phrase"] = "hey tater"
trainer.STATE["training"]["running"] = True
with (
patch.object(trainer, "DATA_DIR", data_dir),
patch.object(trainer, "_ensure_training_venv"),
patch.object(trainer, "_ensure_training_datasets"),
patch.object(trainer.subprocess, "Popen", return_value=process) as popen,
patch.object(trainer, "_normalize_output_artifacts"),
):
trainer._run_training_background(
"hey_tater",
"en",
True,
auto_run=False,
tts_mode="modern",
)
popen.assert_called_once()
log_text = (data_dir / "recorder_training.log").read_text(encoding="utf-8")
self.assertIn("Nvidia Docker Training Run", log_text)
self.assertIn("worker started", log_text)
self.assertFalse(trainer.STATE["training"]["running"])
self.assertEqual(trainer.STATE["training"]["exit_code"], 0)
self.assertIsNone(trainer.TRAINING_THREAD)
finally:
with trainer.STATE_LOCK:
trainer.STATE["raw_phrase"] = original_raw_phrase
trainer.STATE["training"].clear()
trainer.STATE["training"].update(original_training)
trainer.TRAINING_PROCESS = original_process
trainer.TRAINING_THREAD = original_thread
if __name__ == "__main__":
unittest.main()