File size: 5,291 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 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | 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()
|