mirror of
https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git
synced 2026-08-12 07:55:33 -06:00
Release NVIDIA WakeWord Trainer v23
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user