suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
5.29 kB
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()