| from __future__ import annotations |
|
|
| import unittest |
|
|
| from turn_detection.baselines import ( |
| AUDIO_FEATURE_NAMES, |
| LogisticBaseline, |
| extract_audio_statistics, |
| fit_logistic_baseline, |
| fixed_timeout_sweep, |
| ) |
| from turn_detection.runtime.replay import PauseCheckpoint |
|
|
|
|
| class AudioStatisticsBaselineTest(unittest.TestCase): |
| def test_statistics_are_finite_and_have_stable_contract(self) -> None: |
| import numpy as np |
|
|
| sample_rate = 16_000 |
| time = np.arange(sample_rate, dtype=np.float32) / sample_rate |
| audio = 0.2 * np.sin(2 * np.pi * 180 * time) |
| audio[-2_000:] *= np.linspace(1.0, 0.0, 2_000, dtype=np.float32) |
| features = extract_audio_statistics(audio, sample_rate) |
| self.assertEqual(features.shape, (len(AUDIO_FEATURE_NAMES),)) |
| self.assertTrue(np.isfinite(features).all()) |
| self.assertGreaterEqual(features[AUDIO_FEATURE_NAMES.index("periodicity")], 0.0) |
|
|
| def test_logistic_fit_is_deterministic_and_serializable(self) -> None: |
| import numpy as np |
|
|
| negative = np.asarray([[-2.0, -1.0], [-1.5, -0.5], [-1.0, -2.0]]) |
| positive = np.asarray([[2.0, 1.0], [1.5, 0.5], [1.0, 2.0]]) |
| features = np.concatenate((negative, positive), axis=0) |
| labels = np.asarray([0, 0, 0, 1, 1, 1]) |
| first = fit_logistic_baseline(features, labels, epochs=300) |
| second = fit_logistic_baseline(features, labels, epochs=300) |
| self.assertEqual(first, second) |
| restored = LogisticBaseline.from_dict(first.to_dict()) |
| np.testing.assert_allclose(first.predict_proba(features), restored.predict_proba(features)) |
| probabilities = first.predict_proba(features) |
| self.assertTrue((probabilities[:3] < 0.5).all()) |
| self.assertTrue((probabilities[3:] >= 0.5).all()) |
|
|
| def test_fit_requires_both_classes(self) -> None: |
| import numpy as np |
|
|
| with self.assertRaisesRegex(ValueError, "both endpoint classes"): |
| fit_logistic_baseline(np.ones((3, 2)), np.ones(3)) |
|
|
|
|
| class FixedTimeoutBaselineTest(unittest.TestCase): |
| def test_conventional_error_denominators(self) -> None: |
| checkpoints = [ |
| PauseCheckpoint("a", 100.0, 200.0, 0.0, False), |
| PauseCheckpoint("a", 200.0, 600.0, 0.0, True), |
| PauseCheckpoint("b", 100.0, 500.0, 0.0, False), |
| PauseCheckpoint("b", 200.0, 800.0, 0.0, True), |
| ] |
| at_400, at_700 = fixed_timeout_sweep(checkpoints, [400.0, 700.0]) |
| self.assertEqual(at_400["false_positives"], 1) |
| self.assertEqual(at_400["true_negatives"], 1) |
| self.assertEqual(at_400["false_interruption_rate"], 0.5) |
| |
| |
| self.assertEqual(at_400["missed_end_rate"], 0.5) |
| self.assertEqual(at_400["duplicate_response_emissions"], 0) |
| self.assertEqual(at_700["false_interruption_rate"], 0.0) |
| self.assertEqual(at_700["missed_end_rate"], 0.5) |
|
|
| def test_invalid_timeout_is_rejected(self) -> None: |
| with self.assertRaisesRegex(ValueError, "timeouts"): |
| fixed_timeout_sweep([], [-1.0]) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|