| from __future__ import annotations |
|
|
| import math |
| import sys |
| import unittest |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) |
|
|
| from turn_detection.training.metrics import ( |
| average_precision, |
| binary_classification_metrics, |
| confusion_counts, |
| expected_calibration_error, |
| grouped_bootstrap_interval, |
| operational_metrics, |
| roc_auc, |
| sliced_metrics, |
| threshold_at_max_fpr, |
| ) |
|
|
|
|
| class ClassificationMetricsTest(unittest.TestCase): |
| def setUp(self) -> None: |
| self.labels = [0, 0, 1, 1] |
| self.probabilities = [0.1, 0.8, 0.7, 0.9] |
|
|
| def test_confusion_and_conditional_error_denominators(self) -> None: |
| self.assertEqual( |
| confusion_counts(self.labels, self.probabilities, 0.5), |
| {"tp": 2, "fp": 1, "tn": 1, "fn": 0}, |
| ) |
| metrics = binary_classification_metrics(self.labels, self.probabilities, 0.5) |
| self.assertAlmostEqual(metrics["false_positive_rate"], 0.5) |
| self.assertAlmostEqual(metrics["false_negative_rate"], 0.0) |
| self.assertAlmostEqual(metrics["precision"], 2 / 3) |
| self.assertAlmostEqual(metrics["recall"], 1.0) |
| self.assertAlmostEqual(metrics["f1"], 0.8) |
| self.assertAlmostEqual(metrics["accuracy"], 0.75) |
|
|
| def test_rank_metrics_and_ties(self) -> None: |
| self.assertAlmostEqual(roc_auc(self.labels, self.probabilities), 0.75) |
| self.assertAlmostEqual(average_precision(self.labels, self.probabilities), 5 / 6) |
| self.assertAlmostEqual(roc_auc([0, 1], [0.5, 0.5]), 0.5) |
| self.assertAlmostEqual(average_precision([0, 1], [0.5, 0.5]), 0.5) |
|
|
| def test_threshold_maximizes_recall_under_fpr_budget(self) -> None: |
| strict = threshold_at_max_fpr(self.labels, self.probabilities, 0.0) |
| self.assertAlmostEqual(strict["threshold"], 0.9) |
| self.assertAlmostEqual(strict["recall"], 0.5) |
| self.assertAlmostEqual(strict["false_positive_rate"], 0.0) |
| relaxed = threshold_at_max_fpr(self.labels, self.probabilities, 0.5) |
| self.assertAlmostEqual(relaxed["threshold"], 0.7) |
| self.assertAlmostEqual(relaxed["recall"], 1.0) |
|
|
| def test_calibration_is_zero_for_matching_bin_frequencies(self) -> None: |
| |
| self.assertAlmostEqual(expected_calibration_error([0, 1], [0.0, 1.0], 2), 0.0) |
| metrics = binary_classification_metrics([0, 1], [0.0, 1.0]) |
| self.assertAlmostEqual(metrics["brier_score"], 0.0) |
| self.assertLess(metrics["log_loss"], 1e-12) |
|
|
| def test_undefined_conditional_rate_is_none(self) -> None: |
| metrics = binary_classification_metrics([1, 1], [0.8, 0.9]) |
| self.assertIsNone(metrics["false_positive_rate"]) |
| self.assertIsNone(metrics["roc_auc"]) |
|
|
| def test_macro_f1_counts_supported_class_with_no_predictions_as_zero(self) -> None: |
| metrics = binary_classification_metrics([0, 0, 1, 1], [0.1, 0.2, 0.3, 0.4], 0.5) |
| self.assertEqual(metrics["f1"], 0.0) |
| self.assertAlmostEqual(metrics["macro_f1"], 1.0 / 3.0) |
|
|
| def test_rejects_invalid_inputs_instead_of_silently_coercing(self) -> None: |
| with self.assertRaises(ValueError): |
| binary_classification_metrics([0.2, 1.0], [0.2, 0.8]) |
| with self.assertRaises(ValueError): |
| binary_classification_metrics([0, 1], [0.2, math.nan]) |
| with self.assertRaises(ValueError): |
| binary_classification_metrics([], []) |
|
|
|
|
| class ProductMetricsTest(unittest.TestCase): |
| def test_operational_interruptions_are_aggregated_by_turn_and_hour(self) -> None: |
| metrics = operational_metrics( |
| [0, 0, 1, 0], |
| [0.9, 0.1, 0.9, 0.8], |
| threshold=0.5, |
| turn_ids=["a", "a", "a", "b"], |
| total_audio_seconds=3600, |
| endpoint_delays_ms=[100, 200, 300], |
| ) |
| self.assertEqual(metrics["false_interruptions"], 2) |
| self.assertEqual(metrics["turns_with_false_interruption"], 2) |
| self.assertAlmostEqual(metrics["turn_false_interruption_rate"], 1.0) |
| self.assertAlmostEqual(metrics["false_interruptions_per_hour"], 2.0) |
| self.assertAlmostEqual(metrics["endpoint_delay_ms_p50"], 200.0) |
| self.assertAlmostEqual(metrics["endpoint_delay_ms_p95"], 290.0) |
|
|
| def test_slice_metrics(self) -> None: |
| metrics = sliced_metrics( |
| [0, 1, 0, 1], |
| [0.1, 0.9, 0.8, 0.2], |
| {"language": ["hi", "hi", "en", "en"]}, |
| min_count=2, |
| ) |
| self.assertAlmostEqual(metrics["language"]["hi"]["accuracy"], 1.0) |
| self.assertAlmostEqual(metrics["language"]["en"]["accuracy"], 0.0) |
|
|
| def test_group_bootstrap_is_deterministic(self) -> None: |
| first = grouped_bootstrap_interval( |
| [0, 1, 0, 1], |
| [0.2, 0.8, 0.7, 0.3], |
| ["a", "a", "b", "b"], |
| 0.5, |
| samples=50, |
| seed=4, |
| ) |
| second = grouped_bootstrap_interval( |
| [0, 1, 0, 1], |
| [0.2, 0.8, 0.7, 0.3], |
| ["a", "a", "b", "b"], |
| 0.5, |
| samples=50, |
| seed=4, |
| ) |
| self.assertEqual(first, second) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|