tiny-hinglish-turn-detector / tests /test_model_datasets.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
5.82 kB
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()