mirror of
https://github.com/TaterTotterson/microWakeWord-Trainer-Nvidia-Docker.git
synced 2026-08-12 07:55:33 -06:00
Release NVIDIA WakeWord Trainer v22
This commit is contained in:
@@ -292,13 +292,13 @@ class AutoTrainTests(unittest.TestCase):
|
||||
)
|
||||
|
||||
def test_ui_exposes_engine_selector_without_manual_runtime_fields(self):
|
||||
source = (Path(__file__).resolve().parents[1] / "static" / "index.html").read_text(
|
||||
source = (Path(__file__).resolve().parents[1] / "frontend" / "src" / "TrainerApp.vue").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('v-model="trainer.autoForm.stt_engine"', source)
|
||||
self.assertNotIn('trainer.autoForm.stt_model', source)
|
||||
self.assertNotIn('trainer.autoForm.stt_device', source)
|
||||
self.assertNotIn('trainer.autoForm.stt_compute_type', source)
|
||||
self.assertIn("Guided wake check", source)
|
||||
|
||||
def test_phrase_miss_moves_wake_trigger_to_negative_samples(self):
|
||||
@@ -566,6 +566,24 @@ class AutoTrainTests(unittest.TestCase):
|
||||
},
|
||||
)
|
||||
|
||||
def test_trained_word_catalog_keeps_url_alias_for_json_package(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
trained_dir = Path(directory)
|
||||
(trained_dir / "hey_tater.tflite").write_bytes(b"model")
|
||||
(trained_dir / "hey_tater.json").write_text(
|
||||
json.dumps({"wake_word": "hey tater", "model": "hey_tater.tflite"}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with (
|
||||
patch.object(trainer, "TRAINED_WAKE_WORDS_DIR", trained_dir),
|
||||
patch.object(trainer, "_sync_trained_wake_word_artifacts"),
|
||||
):
|
||||
rows = trainer._list_trained_wake_words("http://10.4.20.210:8789")
|
||||
|
||||
self.assertEqual(len(rows), 1)
|
||||
self.assertEqual(rows[0]["url"], rows[0]["json_url"])
|
||||
self.assertTrue(rows[0]["json_url"].endswith("/api/trained_wake_words/hey_tater.json"))
|
||||
|
||||
def test_tater_notification_fails_when_trained_word_is_missing(self):
|
||||
trainer.AUTO_TRAIN_CONFIG["tater_link_token"] = "secret-token"
|
||||
with (
|
||||
|
||||
83
tests/test_data_management.py
Normal file
83
tests/test_data_management.py
Normal file
@@ -0,0 +1,83 @@
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import trainer_server as trainer
|
||||
|
||||
|
||||
class DataManagementTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.tempdir = tempfile.TemporaryDirectory()
|
||||
root = Path(self.tempdir.name)
|
||||
self.original_paths = {
|
||||
"DATA_DIR": trainer.DATA_DIR,
|
||||
"PERSONAL_DIR": trainer.PERSONAL_DIR,
|
||||
"CAPTURED_DIR": trainer.CAPTURED_DIR,
|
||||
"NEGATIVE_DIR": trainer.NEGATIVE_DIR,
|
||||
"TRIM_HISTORY_DIR": trainer.TRIM_HISTORY_DIR,
|
||||
"TRAINED_WAKE_WORDS_DIR": trainer.TRAINED_WAKE_WORDS_DIR,
|
||||
"AUTO_TRAIN_MODEL_DIR": trainer.AUTO_TRAIN_MODEL_DIR,
|
||||
"PIPER_ROOT": trainer.PIPER_ROOT,
|
||||
"PIPER_VOICES_DIR": trainer.PIPER_VOICES_DIR,
|
||||
"PIPER_CATALOG_CACHE_FILE": trainer.PIPER_CATALOG_CACHE_FILE,
|
||||
"OMNIVOICE_CATALOG_CACHE_FILE": trainer.OMNIVOICE_CATALOG_CACHE_FILE,
|
||||
}
|
||||
trainer.DATA_DIR = root
|
||||
trainer.PERSONAL_DIR = root / "personal_samples"
|
||||
trainer.CAPTURED_DIR = root / "captured_audio"
|
||||
trainer.NEGATIVE_DIR = root / "negative_samples"
|
||||
trainer.TRIM_HISTORY_DIR = root / "trim_history"
|
||||
trainer.TRAINED_WAKE_WORDS_DIR = root / "trained_wake_words"
|
||||
trainer.AUTO_TRAIN_MODEL_DIR = root / "auto_train_models"
|
||||
trainer.PIPER_ROOT = root / "tools" / "piper-sample-generator"
|
||||
trainer.PIPER_VOICES_DIR = trainer.PIPER_ROOT / "voices"
|
||||
trainer.PIPER_CATALOG_CACHE_FILE = root / ".cache" / "piper_voices_catalog.json"
|
||||
trainer.OMNIVOICE_CATALOG_CACHE_FILE = root / ".cache" / "omnivoice_languages.json"
|
||||
self.original_training_running = trainer.STATE["training"]["running"]
|
||||
self.original_review_running = trainer.AUTO_TRAIN_RUNTIME["review_running"]
|
||||
trainer.STATE["training"]["running"] = False
|
||||
trainer.AUTO_TRAIN_RUNTIME["review_running"] = False
|
||||
|
||||
def tearDown(self):
|
||||
for name, value in self.original_paths.items():
|
||||
setattr(trainer, name, value)
|
||||
trainer.STATE["training"]["running"] = self.original_training_running
|
||||
trainer.AUTO_TRAIN_RUNTIME["review_running"] = self.original_review_running
|
||||
self.tempdir.cleanup()
|
||||
|
||||
def test_payload_counts_each_managed_item_and_does_not_follow_symlinks(self):
|
||||
generated = trainer.DATA_DIR / "work" / "wake_word_samples"
|
||||
generated.mkdir(parents=True)
|
||||
(generated / "one.wav").write_bytes(b"a" * 128)
|
||||
outside = trainer.DATA_DIR / "outside.bin"
|
||||
outside.write_bytes(b"b" * 8192)
|
||||
(generated / "outside-link").symlink_to(outside)
|
||||
|
||||
payload = trainer._managed_data_payload()
|
||||
item = next(row for row in payload["items"] if row["id"] == "generated_samples")
|
||||
|
||||
self.assertEqual(item["file_count"], 2)
|
||||
self.assertGreater(item["size_bytes"], 0)
|
||||
self.assertEqual(item["location"], "work/wake_word_samples")
|
||||
self.assertEqual(payload["total_file_count"], 2)
|
||||
|
||||
deleted = trainer._delete_managed_data_item("generated_samples")
|
||||
self.assertFalse(generated.exists())
|
||||
self.assertTrue(outside.exists())
|
||||
self.assertEqual(deleted["deleted_id"], "generated_samples")
|
||||
|
||||
def test_unknown_ids_and_active_training_are_rejected(self):
|
||||
with self.assertRaises(KeyError):
|
||||
trainer._delete_managed_data_item("../../not-allowed")
|
||||
|
||||
generated = trainer.DATA_DIR / "work" / "wake_word_samples"
|
||||
generated.mkdir(parents=True)
|
||||
(generated / "keep.wav").write_bytes(b"keep")
|
||||
trainer.STATE["training"]["running"] = True
|
||||
with self.assertRaisesRegex(RuntimeError, "Stop training"):
|
||||
trainer._delete_managed_data_item("generated_samples")
|
||||
self.assertTrue((generated / "keep.wav").exists())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
603
tests/test_modern_tts.py
Normal file
603
tests/test_modern_tts.py
Normal file
@@ -0,0 +1,603 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import importlib.util
|
||||
import json
|
||||
import math
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
import unittest
|
||||
import wave
|
||||
from array import array
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from tts_config import parse_omnivoice_catalog
|
||||
|
||||
try:
|
||||
import trainer_server as trainer
|
||||
except ModuleNotFoundError:
|
||||
trainer = None
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
GENERATOR_PATH = REPO_ROOT / "cli" / "tts_generate_samples.py"
|
||||
SPEC = importlib.util.spec_from_file_location("tts_generate_samples", GENERATOR_PATH)
|
||||
assert SPEC is not None and SPEC.loader is not None
|
||||
generator_module = importlib.util.module_from_spec(SPEC)
|
||||
SPEC.loader.exec_module(generator_module)
|
||||
QA_PATH = REPO_ROOT / "cli" / "tts_reference_qa.py"
|
||||
QA_SPEC = importlib.util.spec_from_file_location("tts_reference_qa", QA_PATH)
|
||||
assert QA_SPEC is not None and QA_SPEC.loader is not None
|
||||
qa_module = importlib.util.module_from_spec(QA_SPEC)
|
||||
QA_SPEC.loader.exec_module(qa_module)
|
||||
|
||||
|
||||
def write_tone(
|
||||
path: Path,
|
||||
*,
|
||||
duration: float = 0.8,
|
||||
amplitude: int = 4000,
|
||||
frequency: float = 220.0,
|
||||
) -> None:
|
||||
rate = 16000
|
||||
samples = array(
|
||||
"h",
|
||||
(
|
||||
int(amplitude * math.sin(2 * math.pi * frequency * index / rate))
|
||||
for index in range(int(rate * duration))
|
||||
),
|
||||
)
|
||||
with wave.open(str(path), "wb") as stream:
|
||||
stream.setnchannels(1)
|
||||
stream.setsampwidth(2)
|
||||
stream.setframerate(rate)
|
||||
stream.writeframes(samples.tobytes())
|
||||
|
||||
|
||||
class ModernTtsTests(unittest.TestCase):
|
||||
def test_direct_generator_uses_one_wake_phrase(self) -> None:
|
||||
self.assertEqual(generator_module.reference_text("hey tater"), "hey tater.")
|
||||
self.assertEqual(generator_module.reference_text("hey tater!"), "hey tater.")
|
||||
self.assertIn("four-provider-direct-corpus", generator_module.GENERATOR_VERSION)
|
||||
self.assertIn("safe-limits", generator_module.GENERATOR_VERSION)
|
||||
|
||||
def test_omnivoice_uses_upstream_sampling_defaults(self) -> None:
|
||||
self.assertEqual(
|
||||
generator_module.omnivoice_stability_args(),
|
||||
["--position_temperature", "5.0", "--class_temperature", "0.0"],
|
||||
)
|
||||
|
||||
def test_omnivoice_uses_a_hidden_stable_prompt_before_short_clone(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=1,
|
||||
batch_size=4,
|
||||
voice_count=2,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
destination = data_dir / "bank"
|
||||
destination.mkdir()
|
||||
|
||||
def create_model_outputs(command, *, only_first: bool = False) -> None:
|
||||
input_flag = "--test_list" if "--test_list" in command else "--input-jsonl"
|
||||
output_flag = "--res_dir" if "--res_dir" in command else "--output-dir"
|
||||
input_path = Path(command[command.index(input_flag) + 1])
|
||||
output_dir = Path(command[command.index(output_flag) + 1])
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
model_entries = [
|
||||
json.loads(line)
|
||||
for line in input_path.read_text(encoding="utf-8").splitlines()
|
||||
if line.strip()
|
||||
]
|
||||
for item in model_entries[:1] if only_first else model_entries:
|
||||
write_tone(output_dir / f"{item['id']}.wav")
|
||||
|
||||
with (
|
||||
patch.object(instance, "ensure_environment"),
|
||||
patch.object(
|
||||
generator_module,
|
||||
"run_with_batch_retry",
|
||||
side_effect=lambda command, _flag, **_kwargs: create_model_outputs(command),
|
||||
) as run_batch,
|
||||
):
|
||||
entries = instance._generate_omni_bank(2, 0, destination)
|
||||
|
||||
self.assertEqual(run_batch.call_count, 2)
|
||||
self.assertEqual(entries[0]["text"], "hey tater.")
|
||||
self.assertEqual(
|
||||
entries[0]["ref_text"],
|
||||
"In a calm and natural voice, I say hey tater clearly, then continue speaking at an even pace.",
|
||||
)
|
||||
self.assertEqual(entries[0]["ref_text"].lower().count("hey tater"), 1)
|
||||
self.assertIn(".omnivoice-prompts", entries[0]["ref_audio"])
|
||||
self.assertEqual(
|
||||
entries[0]["voice_description"],
|
||||
"automatic random voice",
|
||||
)
|
||||
self.assertNotIn("instruct", entries[0])
|
||||
for call in run_batch.call_args_list:
|
||||
command = call.args[0]
|
||||
if "--position_temperature" in command:
|
||||
self.assertEqual(command[command.index("--position_temperature") + 1], "5.0")
|
||||
self.assertEqual(command[command.index("--class_temperature") + 1], "0.0")
|
||||
|
||||
def test_omnivoice_corpus_uses_reference_without_a_second_instruction(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=1,
|
||||
batch_size=4,
|
||||
voice_count=2,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
entries = instance.make_entries(
|
||||
generator_module.ENGINE_OMNIVOICE,
|
||||
1,
|
||||
[{
|
||||
"id": "omni_ref",
|
||||
"path": "/tmp/short.wav",
|
||||
"ref_text": "hey tater.",
|
||||
"omnivoice_prompt_path": "/tmp/prompt.wav",
|
||||
"omnivoice_prompt_text": "A natural carrier sentence.",
|
||||
"instruct": "female, elderly, low pitch, british accent",
|
||||
}],
|
||||
data_dir,
|
||||
)
|
||||
|
||||
self.assertNotIn("instruct", entries[0])
|
||||
|
||||
def test_omnivoice_corpus_repairs_only_vad_rejected_outputs(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=2,
|
||||
batch_size=4,
|
||||
voice_count=2,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
destination = data_dir / "raw"
|
||||
destination.mkdir()
|
||||
entries = [{"id": "omni_a", "text": "hey tater."}, {"id": "omni_b", "text": "hey tater."}]
|
||||
for entry in entries:
|
||||
write_tone(destination / f"{entry['id']}.wav")
|
||||
generation_command = [
|
||||
"omnivoice",
|
||||
"--test_list",
|
||||
str(data_dir / "input.jsonl"),
|
||||
"--res_dir",
|
||||
str(destination),
|
||||
"--batch_size",
|
||||
"4",
|
||||
]
|
||||
qa_calls = 0
|
||||
|
||||
def fake_qa(command, **_kwargs):
|
||||
nonlocal qa_calls
|
||||
qa_calls += 1
|
||||
self.assertIn("--speech-only", command)
|
||||
qa_input = Path(command[command.index("--input-jsonl") + 1])
|
||||
qa_output = Path(command[command.index("--output-jsonl") + 1])
|
||||
candidates = [json.loads(line) for line in qa_input.read_text().splitlines()]
|
||||
results = [
|
||||
{
|
||||
"id": item["id"],
|
||||
"accepted": qa_calls > 1 or item["id"] == "omni_a",
|
||||
}
|
||||
for item in candidates
|
||||
]
|
||||
qa_output.write_text("".join(json.dumps(item) + "\n" for item in results))
|
||||
|
||||
def fake_retry(command, _flag, **_kwargs):
|
||||
retry_input = Path(command[command.index("--test_list") + 1])
|
||||
retry_entries = [json.loads(line) for line in retry_input.read_text().splitlines()]
|
||||
for entry in retry_entries:
|
||||
write_tone(destination / f"{entry['id']}.wav")
|
||||
|
||||
with (
|
||||
patch.object(instance, "_reference_qa_python", return_value=data_dir / "python"),
|
||||
patch.object(generator_module, "run", side_effect=fake_qa),
|
||||
patch.object(generator_module, "run_with_batch_retry", side_effect=fake_retry) as retry,
|
||||
):
|
||||
accepted = instance._repair_generated_corpus(
|
||||
generator_module.ENGINE_OMNIVOICE,
|
||||
entries,
|
||||
destination,
|
||||
generation_command,
|
||||
"",
|
||||
speech_only=True,
|
||||
input_flag="--test_list",
|
||||
batch_flag="--batch_size",
|
||||
)
|
||||
retry_input = Path(retry.call_args.args[0][retry.call_args.args[0].index("--test_list") + 1])
|
||||
retried_ids = [json.loads(line)["id"] for line in retry_input.read_text().splitlines()]
|
||||
|
||||
self.assertEqual([path.name for path in accepted], ["omni_a.wav", "omni_b.wav"])
|
||||
self.assertEqual(retried_ids, ["omni_b"])
|
||||
self.assertEqual(qa_calls, 2)
|
||||
|
||||
def test_omnivoice_repairs_outputs_missing_from_a_successful_seed_batch(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=1,
|
||||
batch_size=4,
|
||||
voice_count=2,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
destination = data_dir / "bank"
|
||||
destination.mkdir()
|
||||
|
||||
def create_outputs(command, *, only_first: bool = False) -> None:
|
||||
input_flag = "--test_list" if "--test_list" in command else "--input-jsonl"
|
||||
output_flag = "--res_dir" if "--res_dir" in command else "--output-dir"
|
||||
input_path = Path(command[command.index(input_flag) + 1])
|
||||
output_dir = Path(command[command.index(output_flag) + 1])
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
model_entries = [
|
||||
json.loads(line)
|
||||
for line in input_path.read_text(encoding="utf-8").splitlines()
|
||||
if line.strip()
|
||||
]
|
||||
for item in model_entries[:1] if only_first else model_entries:
|
||||
write_tone(output_dir / f"{item['id']}.wav")
|
||||
|
||||
batched_calls = 0
|
||||
|
||||
def fake_batched(command, _flag, **_kwargs):
|
||||
nonlocal batched_calls
|
||||
batched_calls += 1
|
||||
create_outputs(command, only_first=batched_calls == 1)
|
||||
|
||||
with (
|
||||
patch.object(instance, "ensure_environment"),
|
||||
patch.object(generator_module, "run_with_batch_retry", side_effect=fake_batched),
|
||||
patch.object(
|
||||
generator_module,
|
||||
"run",
|
||||
side_effect=lambda command, **_kwargs: create_outputs(command),
|
||||
) as run_single,
|
||||
):
|
||||
entries = instance._generate_omni_bank(2, 0, destination)
|
||||
retry_command = run_single.call_args.args[0]
|
||||
retry_input = Path(retry_command[retry_command.index("--test_list") + 1])
|
||||
retried_ids = [json.loads(line)["id"] for line in retry_input.read_text().splitlines()]
|
||||
|
||||
self.assertEqual(len(entries), 2)
|
||||
self.assertEqual(batched_calls, 2)
|
||||
self.assertEqual(run_single.call_count, 1)
|
||||
self.assertEqual(retry_command[retry_command.index("--batch_size") + 1], "1")
|
||||
self.assertEqual(retried_ids, ["omni_prompt_0001"])
|
||||
|
||||
def test_reference_semantic_qa_rejects_noise_and_missing_words(self) -> None:
|
||||
self.assertTrue(qa_module.transcript_matches_phrase("Hey, Tater.", "hey tater"))
|
||||
self.assertTrue(qa_module.transcript_matches_phrase("Hey, gator.", "hey tater"))
|
||||
self.assertFalse(qa_module.transcript_matches_phrase("Tater.", "hey tater"))
|
||||
self.assertFalse(qa_module.transcript_matches_phrase("Hater.", "hey tater"))
|
||||
self.assertFalse(qa_module.transcript_matches_phrase("Thanks for watching!", "hey tater"))
|
||||
self.assertFalse(
|
||||
qa_module.transcript_matches_phrase("Hey tater. Hey tater.", "hey tater")
|
||||
)
|
||||
self.assertFalse(qa_module.transcript_matches_phrase("Hey hey Tate", "hey tater"))
|
||||
self.assertFalse(qa_module.transcript_matches_phrase("", "hey tater"))
|
||||
self.assertEqual(
|
||||
qa_module.semantic_rejection_reason("Ehhhhh...", "hey tater", 0.8),
|
||||
"decoder_collapse",
|
||||
)
|
||||
self.assertEqual(
|
||||
qa_module.semantic_rejection_reason("Hey tater. Hey tater.", "hey tater", 0.8),
|
||||
"repeated_phrase",
|
||||
)
|
||||
self.assertEqual(
|
||||
qa_module.semantic_rejection_reason("Hey hey Tate", "hey tater", 0.8),
|
||||
"repeated_phrase",
|
||||
)
|
||||
self.assertEqual(
|
||||
qa_module.semantic_rejection_reason("Hey Taylor", "hey tater", 0.8),
|
||||
"phrase_mismatch",
|
||||
)
|
||||
|
||||
def test_omnivoice_sample_generation_requires_a_stable_prompt(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=1,
|
||||
batch_size=4,
|
||||
voice_count=2,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
with self.assertRaisesRegex(RuntimeError, "long-form seed prompt"):
|
||||
instance.make_entries(
|
||||
generator_module.ENGINE_OMNIVOICE,
|
||||
1,
|
||||
[{"id": "qwen_ref", "path": "/tmp/qwen.wav", "ref_text": "hey tater."}],
|
||||
data_dir,
|
||||
)
|
||||
|
||||
def test_omnivoice_markdown_catalog_parser(self) -> None:
|
||||
markdown = """
|
||||
| # | Language | OmniVoice ID | ISO 639-3 | Duration (h) |
|
||||
|--:|----------|:------------:|:---------:|:------------:|
|
||||
| 1 | English | en | eng | 100000.5 |
|
||||
| 2 | Amdo Tibetan | adx | adx | 56.94 |
|
||||
"""
|
||||
parsed = parse_omnivoice_catalog(markdown)
|
||||
self.assertEqual(parsed["en"]["name"], "English")
|
||||
self.assertEqual(parsed["adx"]["iso_639_3"], "adx")
|
||||
self.assertEqual(parsed["adx"]["duration_hours"], 56.94)
|
||||
|
||||
@unittest.skipIf(trainer is None, "trainer server dependencies are not installed")
|
||||
def test_language_catalog_merges_engine_coverage_and_quality(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
with (
|
||||
patch.object(
|
||||
trainer,
|
||||
"_load_omnivoice_catalog",
|
||||
return_value={
|
||||
"en": {"name": "English"},
|
||||
"zu": {"name": "Zulu"},
|
||||
},
|
||||
),
|
||||
patch.object(trainer, "_load_piper_catalog", return_value={}),
|
||||
patch.object(trainer, "PIPER_ROOT", Path(temp_dir) / "piper"),
|
||||
patch.object(trainer, "PIPER_VOICES_DIR", Path(temp_dir) / "voices"),
|
||||
):
|
||||
catalog = {item["code"]: item for item in trainer._available_languages()}
|
||||
|
||||
self.assertEqual(catalog["en"]["quality"], "recommended")
|
||||
self.assertEqual(catalog["en"]["engines"], ["omnivoice", "qwen3", "moss"])
|
||||
self.assertEqual(catalog["zu"]["quality"], "experimental")
|
||||
self.assertEqual(catalog["zu"]["engines"], ["omnivoice"])
|
||||
|
||||
@unittest.skipIf(trainer is None, "trainer server dependencies are not installed")
|
||||
def test_server_resolves_unavailable_tts_modes_safely(self) -> None:
|
||||
languages = [
|
||||
{"code": "en", "engines": ["omnivoice", "qwen3", "moss"]},
|
||||
{"code": "legacy", "engines": ["piper"]},
|
||||
]
|
||||
self.assertEqual(
|
||||
trainer._resolve_tts_mode_for_language("piper", "en", languages),
|
||||
"modern",
|
||||
)
|
||||
self.assertEqual(
|
||||
trainer._resolve_tts_mode_for_language("modern", "legacy", languages),
|
||||
"piper",
|
||||
)
|
||||
|
||||
def test_generator_plan_and_piper_discovery_do_not_load_models(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
output_dir = data_dir / "work" / "wake_word_samples"
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=101,
|
||||
batch_size=8,
|
||||
voice_count=128,
|
||||
data_dir=data_dir,
|
||||
output_dir=output_dir,
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=True,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
self.assertEqual(instance.spoken_phrase, "hey tater")
|
||||
self.assertEqual(instance.reference_text, "hey tater.")
|
||||
self.assertEqual(instance.voice_bank_dir.name, generator_module.phrase_key("hey tater"))
|
||||
self.assertEqual(instance.engines(), ["omnivoice", "qwen3", "moss"])
|
||||
self.assertEqual(sum(generator_module.distribute_samples(101, instance.engines()).values()), 101)
|
||||
|
||||
model = data_dir / "tools" / "piper-sample-generator" / "models" / "en_US-libritts_r-medium.pt"
|
||||
model.parent.mkdir(parents=True)
|
||||
model.touch()
|
||||
args.tts_mode = "hybrid"
|
||||
self.assertEqual(instance.engines()[-1], "piper")
|
||||
|
||||
def test_voice_descriptions_are_distinct_for_default_bank(self) -> None:
|
||||
descriptions = generator_module.qwen_descriptions("English", 128)
|
||||
self.assertEqual(len(descriptions), 128)
|
||||
self.assertEqual(len(set(descriptions)), 128)
|
||||
|
||||
first_bank = descriptions[:64]
|
||||
self.assertEqual(sum(" female speaker " in item for item in first_bank), 32)
|
||||
self.assertEqual(sum(" male speaker " in item for item in first_bank), 32)
|
||||
for trait in (
|
||||
"child",
|
||||
"teenager",
|
||||
"young adult",
|
||||
"middle-aged adult",
|
||||
"elderly adult",
|
||||
"low pitch",
|
||||
"medium pitch",
|
||||
"high pitch",
|
||||
"calm neutral delivery",
|
||||
"bright energetic delivery",
|
||||
"soft careful delivery",
|
||||
"confident resonant delivery",
|
||||
"casual conversational delivery",
|
||||
"clear timbre",
|
||||
"warm timbre",
|
||||
"slightly breathy timbre",
|
||||
"crisp timbre",
|
||||
"gently rough timbre",
|
||||
):
|
||||
self.assertGreaterEqual(sum(trait in item for item in first_bank), 10, trait)
|
||||
|
||||
def test_failed_model_batch_retries_one_item_at_a_time(self) -> None:
|
||||
command = ["worker", "--batch-size", "4"]
|
||||
with patch.object(
|
||||
generator_module,
|
||||
"run",
|
||||
side_effect=(subprocess.CalledProcessError(1, command), None),
|
||||
) as mocked_run:
|
||||
generator_module.run_with_batch_retry(command, "--batch-size")
|
||||
|
||||
self.assertEqual(mocked_run.call_count, 2)
|
||||
self.assertEqual(mocked_run.call_args_list[1].args[0], ["worker", "--batch-size", "1"])
|
||||
|
||||
def test_acoustic_qa_accepts_speech_like_pcm_and_rejects_silence(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
root = Path(temp_dir)
|
||||
tone = root / "tone.wav"
|
||||
silence = root / "silence.wav"
|
||||
write_tone(tone)
|
||||
write_tone(silence, amplitude=0)
|
||||
self.assertTrue(generator_module.valid_sample(tone))
|
||||
self.assertTrue(generator_module.valid_reference(tone))
|
||||
self.assertFalse(generator_module.valid_sample(silence))
|
||||
|
||||
def test_provider_safety_gate_rejects_static_and_rambling(self) -> None:
|
||||
clean = {
|
||||
"duration": 1.2,
|
||||
"rms": 0.08,
|
||||
"peak": 0.5,
|
||||
"clipped_ratio": 0.0,
|
||||
"dc_offset": 0.0,
|
||||
"spectral_flatness": 0.05,
|
||||
"high_frequency_ratio": 0.04,
|
||||
"zero_crossing_rate": 0.08,
|
||||
}
|
||||
self.assertEqual(
|
||||
qa_module.acoustic_rejection_reason(clean, 0.7, "omnivoice", 0.4, 2.7),
|
||||
"accepted",
|
||||
)
|
||||
self.assertEqual(
|
||||
qa_module.acoustic_rejection_reason(
|
||||
{**clean, "spectral_flatness": 0.8}, 0.8, "omnivoice", 0.4, 2.7
|
||||
),
|
||||
"static_or_broadband_noise",
|
||||
)
|
||||
self.assertEqual(
|
||||
qa_module.acoustic_rejection_reason(
|
||||
{**clean, "duration": 3.5}, 0.8, "qwen3", 0.4, 2.7
|
||||
),
|
||||
"too_long_or_rambling",
|
||||
)
|
||||
|
||||
def test_direct_entries_do_not_clone_the_old_voice_bank(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
args = argparse.Namespace(
|
||||
phrase="hey_tater",
|
||||
language="en",
|
||||
tts_mode="hybrid",
|
||||
samples=12,
|
||||
batch_size=4,
|
||||
voice_count=128,
|
||||
data_dir=data_dir,
|
||||
output_dir=data_dir / "work" / "samples",
|
||||
ffmpeg="ffmpeg",
|
||||
dry_run=False,
|
||||
)
|
||||
instance = generator_module.Generator(args)
|
||||
qwen = instance.make_direct_entries("qwen3", 4, data_dir, [])
|
||||
omni = instance.make_direct_entries("omnivoice", 4, data_dir, [])
|
||||
refs = [data_dir / f"accepted-{index}.wav" for index in range(4)]
|
||||
moss = instance.make_direct_entries("moss", 4, data_dir, refs)
|
||||
|
||||
self.assertTrue(all("ref_audio" not in item for item in qwen + omni))
|
||||
self.assertEqual(len({item["instruct"] for item in qwen}), 4)
|
||||
self.assertEqual([item["ref_audio"] for item in moss], [str(path) for path in refs])
|
||||
|
||||
@unittest.skipUnless(shutil.which("ffmpeg"), "ffmpeg is required for normalization")
|
||||
def test_orchestrator_produces_exact_normalized_corpus_and_manifest(self) -> None:
|
||||
class FakeGenerator(generator_module.Generator):
|
||||
generated = 0
|
||||
|
||||
def generate_direct_engine(self, engine, count, reference_paths, prefix=""):
|
||||
destination = self.raw_dir / f"{engine}_{prefix or 'main'}"
|
||||
destination.mkdir(parents=True, exist_ok=True)
|
||||
paths = []
|
||||
for index in range(count):
|
||||
path = destination / f"{engine}_{prefix}{index}.wav"
|
||||
write_tone(path, frequency=180 + self.generated)
|
||||
self.generated += 1
|
||||
self.speed_by_path[path.resolve()] = 1.0
|
||||
paths.append(path)
|
||||
entries = [
|
||||
{
|
||||
"id": path.stem,
|
||||
"minimum_duration": self.minimum_duration,
|
||||
"maximum_duration": self.maximum_duration,
|
||||
}
|
||||
for path in paths
|
||||
]
|
||||
return entries, paths
|
||||
|
||||
def qualify_direct_candidates(self, engine, entries, paths, prefix=""):
|
||||
return paths
|
||||
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
data_dir = Path(temp_dir)
|
||||
output_dir = data_dir / "work" / "wake_word_samples"
|
||||
args = argparse.Namespace(
|
||||
phrase="hey tater",
|
||||
language="en",
|
||||
tts_mode="modern",
|
||||
samples=13,
|
||||
batch_size=4,
|
||||
voice_count=8,
|
||||
data_dir=data_dir,
|
||||
output_dir=output_dir,
|
||||
ffmpeg=shutil.which("ffmpeg"),
|
||||
dry_run=False,
|
||||
)
|
||||
instance = FakeGenerator(args)
|
||||
instance.generate()
|
||||
|
||||
self.assertEqual(len(list(output_dir.glob("*.wav"))), 13)
|
||||
self.assertTrue((output_dir / ".generation_manifest.json").is_file())
|
||||
self.assertTrue(instance.cache_hit())
|
||||
|
||||
def test_docker_and_ui_are_wired_for_modern_tts(self) -> None:
|
||||
for dockerfile in ("dockerfile", "dockerfile.blackwell"):
|
||||
source = (REPO_ROOT / dockerfile).read_text(encoding="utf-8")
|
||||
self.assertIn("ffmpeg", source)
|
||||
self.assertIn("tts_config.py", source)
|
||||
|
||||
ui = (REPO_ROOT / "frontend" / "src" / "TrainerApp.vue").read_text(encoding="utf-8")
|
||||
store = (REPO_ROOT / "frontend" / "src" / "trainerStore.ts").read_text(encoding="utf-8")
|
||||
self.assertIn('v-model="trainer.ttsMode"', ui)
|
||||
self.assertIn("tts_mode: trainer.ttsMode", store)
|
||||
self.assertIn("OmniVoice", store)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
54
tests/test_session_stop.py
Normal file
54
tests/test_session_stop.py
Normal file
@@ -0,0 +1,54 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import signal
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
import trainer_server as trainer
|
||||
|
||||
|
||||
class _FakeTrainingProcess:
|
||||
def __init__(self):
|
||||
self.pid = 5432
|
||||
self.returncode = None
|
||||
|
||||
def poll(self):
|
||||
return self.returncode
|
||||
|
||||
def wait(self, timeout=None):
|
||||
self.returncode = -signal.SIGTERM
|
||||
return self.returncode
|
||||
|
||||
def terminate(self):
|
||||
self.returncode = -signal.SIGTERM
|
||||
|
||||
def kill(self):
|
||||
self.returncode = -signal.SIGKILL
|
||||
|
||||
|
||||
class SessionStopTests(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
trainer.TRAINING_STOP_EVENT.clear()
|
||||
|
||||
def test_session_stop_terminates_the_process_group_and_allows_another_run(self):
|
||||
proc = _FakeTrainingProcess()
|
||||
original_process = trainer.TRAINING_PROCESS
|
||||
original_thread = trainer.TRAINING_THREAD
|
||||
try:
|
||||
trainer.TRAINING_PROCESS = proc
|
||||
trainer.TRAINING_THREAD = None
|
||||
with (
|
||||
patch.object(trainer.os, "getpgid", return_value=proc.pid),
|
||||
patch.object(trainer.os, "getpgrp", return_value=999),
|
||||
patch.object(trainer.os, "killpg") as killpg,
|
||||
):
|
||||
self.assertTrue(trainer._stop_current_training(timeout=0.2))
|
||||
killpg.assert_called_once_with(proc.pid, signal.SIGTERM)
|
||||
self.assertFalse(trainer.TRAINING_STOP_EVENT.is_set())
|
||||
finally:
|
||||
trainer.TRAINING_PROCESS = original_process
|
||||
trainer.TRAINING_THREAD = original_thread
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
61
tests/test_tts_config.py
Normal file
61
tests/test_tts_config.py
Normal file
@@ -0,0 +1,61 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from tts_config import (
|
||||
ENGINE_MOSS,
|
||||
ENGINE_OMNIVOICE,
|
||||
ENGINE_PIPER,
|
||||
ENGINE_QWEN3,
|
||||
distribute_samples,
|
||||
engines_for_language,
|
||||
language_for_engine,
|
||||
normalize_tts_mode,
|
||||
quality_for_engines,
|
||||
)
|
||||
|
||||
|
||||
class TtsConfigTests(unittest.TestCase):
|
||||
def test_recommended_languages_use_all_modern_engines(self) -> None:
|
||||
self.assertEqual(
|
||||
engines_for_language("en", "modern"),
|
||||
[ENGINE_OMNIVOICE, ENGINE_QWEN3, ENGINE_MOSS],
|
||||
)
|
||||
self.assertEqual(
|
||||
quality_for_engines(engines_for_language("fr", "modern")),
|
||||
"recommended",
|
||||
)
|
||||
|
||||
def test_broad_language_coverage_routes_through_omnivoice(self) -> None:
|
||||
self.assertEqual(engines_for_language("zu", "modern"), [ENGINE_OMNIVOICE])
|
||||
self.assertEqual(quality_for_engines([ENGINE_OMNIVOICE]), "experimental")
|
||||
|
||||
def test_hybrid_and_legacy_modes_require_available_piper(self) -> None:
|
||||
self.assertEqual(
|
||||
engines_for_language("en", "hybrid", piper_available=True),
|
||||
[ENGINE_OMNIVOICE, ENGINE_QWEN3, ENGINE_MOSS, ENGINE_PIPER],
|
||||
)
|
||||
self.assertEqual(engines_for_language("en", "piper"), [])
|
||||
self.assertEqual(
|
||||
engines_for_language("en", "piper", piper_available=True),
|
||||
[ENGINE_PIPER],
|
||||
)
|
||||
|
||||
def test_sample_distribution_is_exact_and_deterministic(self) -> None:
|
||||
self.assertEqual(
|
||||
distribute_samples(10, [ENGINE_OMNIVOICE, ENGINE_QWEN3, ENGINE_MOSS]),
|
||||
{ENGINE_OMNIVOICE: 4, ENGINE_QWEN3: 3, ENGINE_MOSS: 3},
|
||||
)
|
||||
self.assertEqual(sum(distribute_samples(50000, ["a", "b", "c"]).values()), 50000)
|
||||
|
||||
def test_invalid_mode_falls_back_to_four_provider_route(self) -> None:
|
||||
self.assertEqual(normalize_tts_mode("unknown"), "hybrid")
|
||||
|
||||
def test_common_language_aliases_use_model_catalog_ids(self) -> None:
|
||||
self.assertEqual(language_for_engine(ENGINE_OMNIVOICE, "ar"), "arb")
|
||||
self.assertEqual(language_for_engine(ENGINE_OMNIVOICE, "ne"), "npi")
|
||||
self.assertEqual(language_for_engine(ENGINE_MOSS, "ar"), "ar")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
109
tests/test_vue_ui.py
Normal file
109
tests/test_vue_ui.py
Normal file
@@ -0,0 +1,109 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import pathlib
|
||||
import unittest
|
||||
|
||||
|
||||
REPO_ROOT = pathlib.Path(__file__).resolve().parents[1]
|
||||
|
||||
|
||||
class VueTrainerUiTests(unittest.TestCase):
|
||||
def test_frontend_uses_typed_vue_and_vite(self) -> None:
|
||||
package = json.loads((REPO_ROOT / "frontend" / "package.json").read_text(encoding="utf-8"))
|
||||
self.assertEqual(package["dependencies"]["vue"], "3.5.40")
|
||||
self.assertIn("vue-tsc --noEmit", package["scripts"]["build"])
|
||||
self.assertIn("vite build", package["scripts"]["build"])
|
||||
|
||||
config = (REPO_ROOT / "frontend" / "vite.config.ts").read_text(encoding="utf-8")
|
||||
self.assertIn('"../static/ui"', config)
|
||||
self.assertIn('fileName: () => "trainer-ui.js"', config)
|
||||
|
||||
def test_static_shell_loads_prebuilt_bundle(self) -> None:
|
||||
index = (REPO_ROOT / "static" / "index.html").read_text(encoding="utf-8")
|
||||
self.assertIn('id="trainer-app"', index)
|
||||
self.assertIn('/static/ui/trainer-ui.css', index)
|
||||
self.assertIn('/static/ui/trainer-ui.js', index)
|
||||
self.assertNotIn("fonts.googleapis.com", index)
|
||||
|
||||
self.assertGreater((REPO_ROOT / "static" / "ui" / "trainer-ui.js").stat().st_size, 100_000)
|
||||
self.assertGreater((REPO_ROOT / "static" / "ui" / "trainer-ui.css").stat().st_size, 10_000)
|
||||
|
||||
def test_theme_uses_tater_orange_and_neutral_greys(self) -> None:
|
||||
styles = (REPO_ROOT / "frontend" / "src" / "trainer.css").read_text(encoding="utf-8")
|
||||
self.assertIn("--orange: #ff9134", styles)
|
||||
self.assertIn("--surface: rgba(29, 29, 31, .9)", styles)
|
||||
for old_blue in ("#070b15", "#11192b", "#5db6ff", "#8d75ff", "#7fc7ff", "#78caff"):
|
||||
self.assertNotIn(old_blue, styles)
|
||||
|
||||
def test_reactive_ui_keeps_trainer_workflows(self) -> None:
|
||||
app = (REPO_ROOT / "frontend" / "src" / "TrainerApp.vue").read_text(encoding="utf-8")
|
||||
store = (REPO_ROOT / "frontend" / "src" / "trainerStore.ts").read_text(encoding="utf-8")
|
||||
trim = (REPO_ROOT / "frontend" / "src" / "components" / "AudioTrimModal.vue").read_text(encoding="utf-8")
|
||||
|
||||
for workflow in (
|
||||
"startSession",
|
||||
"stopSession",
|
||||
"startTraining",
|
||||
"saveAuto",
|
||||
"runAutoAction",
|
||||
"reviewCaptured",
|
||||
"uploadSelectedFiles",
|
||||
"copyWakeWord",
|
||||
"deleteManagedData",
|
||||
):
|
||||
self.assertIn(workflow, app)
|
||||
|
||||
for endpoint in (
|
||||
"/api/start_session",
|
||||
"/api/stop_session",
|
||||
"/api/upload_personal_sample",
|
||||
"/api/captured_audio",
|
||||
"/api/auto_train",
|
||||
"/api/train_status",
|
||||
"/api/trained_wake_words/catalog",
|
||||
"/api/data",
|
||||
):
|
||||
self.assertIn(endpoint, store)
|
||||
|
||||
self.assertIn("OfflineAudioContext", trim)
|
||||
self.assertIn("/api/samples/trim", trim)
|
||||
self.assertIn(':disabled="Boolean(trainer.session.safe_word)', app)
|
||||
self.assertIn('{ id: "data", label: "Data"', app)
|
||||
|
||||
def test_training_console_pauses_follow_mode_when_scrolled_up(self) -> None:
|
||||
app = (REPO_ROOT / "frontend" / "src" / "TrainerApp.vue").read_text(encoding="utf-8")
|
||||
|
||||
self.assertIn("const consoleFollowing = ref(true)", app)
|
||||
self.assertIn("distanceFromBottom <= 32", app)
|
||||
self.assertIn('if (!consoleFollowing.value) return', app)
|
||||
self.assertIn('@scroll.passive="onConsoleScroll"', app)
|
||||
self.assertIn("Jump to latest", app)
|
||||
|
||||
def test_wake_word_card_uses_explicit_json_catalog_url(self) -> None:
|
||||
app = (REPO_ROOT / "frontend" / "src" / "TrainerApp.vue").read_text(encoding="utf-8")
|
||||
types = (REPO_ROOT / "frontend" / "src" / "types.ts").read_text(encoding="utf-8")
|
||||
|
||||
self.assertIn("item.json_url || item.url || item.jsonUrl", app)
|
||||
self.assertIn("copyWakeWord(wordJsonUrl(word))", app)
|
||||
self.assertNotIn("copyWakeWord(word.url)", app)
|
||||
self.assertIn("json_url?: string", types)
|
||||
|
||||
def test_runtime_packaging_uses_bundle_without_node(self) -> None:
|
||||
dockerfiles = [REPO_ROOT / "dockerfile", REPO_ROOT / "dockerfile.blackwell"]
|
||||
for dockerfile in dockerfiles:
|
||||
if not dockerfile.exists():
|
||||
continue
|
||||
source = dockerfile.read_text(encoding="utf-8")
|
||||
self.assertIn("COPY --chown=root:root static/ /root/mww-scripts/static/", source)
|
||||
self.assertNotIn("npm install", source)
|
||||
|
||||
macos_builder = REPO_ROOT / "macos" / "WakeWordTrainer" / "scripts" / "build_app.sh"
|
||||
if macos_builder.exists():
|
||||
source = macos_builder.read_text(encoding="utf-8")
|
||||
self.assertIn("--exclude='frontend/node_modules/'", source)
|
||||
self.assertNotIn("--exclude='static/'", source)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user