| 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: |
| 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() |
|
|