tiny-hinglish-turn-detector / tests /test_model_tiny_tcn.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
6.23 kB
from __future__ import annotations
import sys
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 - lightweight CI
torch = None
@unittest.skipUnless(torch is not None, "PyTorch is not installed")
class TinyTCNTest(unittest.TestCase):
def setUp(self) -> None:
from turn_detection.models.tiny_tcn import TinyTCNConfig, TinyTurnDetector
self.config = TinyTCNConfig(
channels=32,
num_blocks=3,
kernel_size=3,
dilation_cycle=(1, 2, 4),
attention_channels=16,
head_hidden=16,
dropout=0.0,
)
torch.manual_seed(3)
self.model = TinyTurnDetector(self.config).eval()
def test_output_shapes_are_logits(self) -> None:
features = torch.randn(3, 80, 51)
mask = torch.ones(3, 51, dtype=torch.bool)
output = self.model(features, mask)
self.assertEqual(tuple(output.endpoint_logits.shape), (3,))
self.assertEqual(tuple(output.midfiller_logits.shape), (3,))
self.assertEqual(tuple(output.endfiller_logits.shape), (3,))
self.assertTrue(bool(torch.isfinite(output.endpoint_logits).all()))
def test_right_padding_does_not_change_valid_prediction(self) -> None:
features = torch.randn(2, 80, 40)
mask = torch.ones(2, 40, dtype=torch.bool)
padded = torch.cat([features, torch.randn(2, 80, 17) * 100], dim=-1)
padded_mask = torch.cat([mask, torch.zeros(2, 17, dtype=torch.bool)], dim=-1)
with torch.inference_mode():
original = self.model(features, mask).endpoint_logits
after_padding = self.model(padded, padded_mask).endpoint_logits
torch.testing.assert_close(original, after_padding, atol=1e-6, rtol=1e-6)
def test_default_model_is_in_intended_tiny_parameter_range(self) -> None:
from turn_detection.models.tiny_tcn import TinyTurnDetector
count = sum(parameter.numel() for parameter in TinyTurnDetector().parameters())
self.assertGreaterEqual(count, 300_000)
self.assertLess(count, 1_000_000)
def test_invalid_feature_shape_is_rejected(self) -> None:
with self.assertRaises(ValueError):
self.model(torch.randn(2, 79, 20))
@unittest.skipUnless(torch is not None, "PyTorch is not installed")
class FrontendTest(unittest.TestCase):
def test_feature_and_mask_lengths(self) -> None:
from turn_detection.models.features import LogMelFrontend
frontend = LogMelFrontend()
waveform = torch.randn(2, 16_000)
features, mask = frontend(waveform, torch.tensor([16_000, 8_000]))
self.assertEqual(tuple(features.shape[:2]), (2, 80))
self.assertEqual(int(mask[0].sum()), 98)
self.assertEqual(int(mask[1].sum()), 48)
self.assertTrue(bool((features[1, :, ~mask[1]] == 0).all()))
def test_deployment_frontend_parity(self) -> None:
try:
import numpy as np
except ImportError:
self.skipTest("numpy is not installed")
from turn_detection.models.features import LogMelConfig, LogMelFrontend
from turn_detection.runtime.features import FrontendConfig, log_mel_spectrogram
audio = np.random.default_rng(17).standard_normal(12_345).astype(np.float32) * 0.1
runtime_features, runtime_mask = log_mel_spectrogram(
audio,
16_000,
FrontendConfig(max_seconds=1.0, normalization="whisper", pad_side="left"),
)
training_frontend = LogMelFrontend(
LogMelConfig(
normalize=False,
mel_scale="htk",
log_scale="whisper",
center=True,
drop_last_frame=True,
pad_side="left",
)
)
padded = torch.zeros(1, 16_000)
padded[0, -len(audio) :] = torch.from_numpy(audio)
training_features, training_mask = training_frontend(padded, torch.tensor([len(audio)]))
torch.testing.assert_close(
training_features[0],
torch.from_numpy(runtime_features),
atol=5e-6,
rtol=1e-5,
)
self.assertEqual(training_mask[0].to(torch.float32).tolist(), runtime_mask.tolist())
def test_export_metadata_loads_in_runtime(self) -> None:
import json
import tempfile
from turn_detection.models.deployment import build_runtime_metadata
from turn_detection.models.features import LogMelConfig
from turn_detection.runtime.predictor import ModelMetadata
config = LogMelConfig(
normalize=False,
log_scale="whisper",
center=True,
drop_last_frame=True,
pad_side="left",
)
payload = build_runtime_metadata(
config,
max_seconds=4.0,
threshold=0.61,
model_name="preview",
architecture="tiny_tcn",
development_only=True,
training_status="preview-only",
data_scope="one shard",
data_revision="abc123",
parameter_count=151_812,
)
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "model_metadata.json"
path.write_text(json.dumps(payload), encoding="utf-8")
loaded = ModelMetadata.from_path(path)
self.assertEqual(loaded.input_features_name, "log_mel")
self.assertEqual(loaded.frame_mask_name, "frame_mask")
self.assertEqual(loaded.output_type, "probability")
self.assertEqual(loaded.frontend.target_frames, 400)
self.assertTrue(loaded.development_only)
self.assertEqual(loaded.training_status, "preview-only")
self.assertEqual(loaded.data_scope, "one shard")
self.assertEqual(loaded.data_revision, "abc123")
self.assertEqual(loaded.parameter_count, 151_812)
self.assertEqual(loaded.controller.endpoint_threshold, 0.61)
self.assertAlmostEqual(loaded.controller.long_pause_threshold, 0.43)
if __name__ == "__main__":
unittest.main()