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