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