File size: 5,820 Bytes
35d483e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | from __future__ import annotations
import json
import sys
import tempfile
import unittest
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
try:
import torch
except ImportError: # pragma: no cover
torch = None
@unittest.skipUnless(torch is not None, "PyTorch is not installed")
class ManifestRoutingTest(unittest.TestCase):
def test_collator_distinguishes_observed_turn_ids_from_record_fallbacks(self) -> None:
from turn_detection.models.features import LogMelFrontend
from turn_detection.training.datasets import AudioFeatureCollator
collator = AudioFeatureCollator(LogMelFrontend(), max_seconds=1.0)
batch = collator(
[
{
"record_id": "clip-only",
"log_mel": torch.zeros(80, 4),
"endpoint": False,
},
{
"record_id": "clip-in-turn",
"turn_id": "turn-7",
"log_mel": torch.zeros(80, 4),
"endpoint": True,
},
]
)
self.assertEqual(batch["turn_id"], ["clip-only", "turn-7"])
self.assertEqual(batch["turn_id_observed"], [False, True])
def test_nonexistent_audio_basename_uses_parquet_resolver(self) -> None:
from turn_detection.models.features import LogMelFrontend
from turn_detection.training.datasets import (
ManifestAudioStream,
build_record_dataloader,
)
with tempfile.TemporaryDirectory() as directory:
manifest = Path(directory) / "split.jsonl"
manifest.write_text(
json.dumps(
{
"record_id": "one",
"split": "train",
"audio_path": "original-basename.flac",
"source_file": "data-00000.parquet",
"source_row": 0,
"endpoint": True,
}
)
+ "\n",
encoding="utf-8",
)
loader = build_record_dataloader(
manifest,
split="train",
frontend=LogMelFrontend(),
batch_size=1,
max_seconds=8.0,
shuffle=False,
)
self.assertIsInstance(loader.dataset, ManifestAudioStream)
def test_existing_audio_path_remains_direct(self) -> None:
from turn_detection.models.features import LogMelFrontend
from turn_detection.training.datasets import (
ManifestDataset,
build_record_dataloader,
)
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "clip.wav").touch()
manifest = root / "split.jsonl"
manifest.write_text(
json.dumps(
{
"record_id": "one",
"split": "train",
"audio_path": "clip.wav",
"source_file": "data-00000.parquet",
"source_row": 0,
"endpoint": True,
}
)
+ "\n",
encoding="utf-8",
)
loader = build_record_dataloader(
manifest,
split="train",
frontend=LogMelFrontend(),
batch_size=1,
max_seconds=8.0,
shuffle=False,
)
self.assertIsInstance(loader.dataset, ManifestDataset)
@unittest.skipUnless(torch is not None, "PyTorch is not installed")
class ResamplingParityTest(unittest.TestCase):
def test_8khz_and_48khz_match_runtime_contract(self) -> None:
try:
import numpy as np
except ImportError:
self.skipTest("numpy is not installed")
from turn_detection.runtime.features import resample_waveform
from turn_detection.training.datasets import decode_audio
for sample_rate in (8_000, 48_000):
with self.subTest(sample_rate=sample_rate):
samples = (
np.random.default_rng(sample_rate)
.standard_normal(sample_rate // 3)
.astype("float32")
* 0.1
)
expected = resample_waveform(samples, sample_rate, 16_000)
actual = decode_audio(
{"audio": samples, "sample_rate": sample_rate},
target_sample_rate=16_000,
max_seconds=2.0,
)
np.testing.assert_allclose(actual.numpy(), expected, atol=1e-7, rtol=1e-6)
def test_long_audio_is_cropped_before_resampling_in_both_paths(self) -> None:
try:
import numpy as np
except ImportError:
self.skipTest("numpy is not installed")
from turn_detection.runtime.features import resample_waveform
from turn_detection.training.datasets import decode_audio
sample_rate = 48_000
max_seconds = 2.0
samples = (
np.random.default_rng(19).standard_normal(sample_rate * 30).astype("float32") * 0.1
)
suffix = samples[-round(sample_rate * max_seconds) :]
expected = resample_waveform(suffix, sample_rate, 16_000)
actual = decode_audio(
{"audio": samples, "sample_rate": sample_rate},
target_sample_rate=16_000,
max_seconds=max_seconds,
)
np.testing.assert_allclose(actual.numpy(), expected, atol=1e-7, rtol=1e-6)
if __name__ == "__main__":
unittest.main()
|