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) # Turn b already emitted a premature response at 500 ms; the later true # endpoint cannot emit again within the same latched turn. 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()