File size: 3,266 Bytes
35d483e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
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()