tiny-hinglish-turn-detector / tests /test_baselines.py
suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
3.27 kB
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()