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