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 ( # noqa: E402 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: # Each half-width bin has observed rate equal to its mean confidence. 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()