85 lines
3.1 KiB
Python
85 lines
3.1 KiB
Python
"""Stdlib tests for eval metrics + the permutation test.
|
|
|
|
Must not open Chroma — the template has no corpus. Run with:
|
|
|
|
python -m unittest eval.test_metrics
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import math
|
|
import unittest
|
|
|
|
from eval.pvalue import compare, paired_permutation
|
|
from eval.run_eval import ndcg_at_k, p_at_1, recall_at_k, reciprocal_rank
|
|
|
|
|
|
A, B, C, X, Y = ("b", "a"), ("b", "b"), ("b", "c"), ("b", "x"), ("b", "y")
|
|
|
|
|
|
class MetricTests(unittest.TestCase):
|
|
def test_reciprocal_rank(self) -> None:
|
|
# Q1: expected at rank 1
|
|
self.assertEqual(reciprocal_rank([A, B, C], [A]), 1.0)
|
|
# Q2: expected at rank 2
|
|
self.assertEqual(reciprocal_rank([B, A], [A]), 0.5)
|
|
# Q3: miss
|
|
self.assertEqual(reciprocal_rank([X, Y], [A]), 0.0)
|
|
|
|
def test_p_at_1(self) -> None:
|
|
self.assertEqual(p_at_1([A, B], [A]), 1.0)
|
|
self.assertEqual(p_at_1([B, A], [A]), 0.0)
|
|
self.assertEqual(p_at_1([], [A]), 0.0)
|
|
self.assertEqual(p_at_1([A], []), 0.0)
|
|
|
|
def test_recall_at_k(self) -> None:
|
|
self.assertEqual(recall_at_k([A, B, C], [A], 1), 1.0)
|
|
self.assertEqual(recall_at_k([B, A], [A], 1), 0.0)
|
|
self.assertEqual(recall_at_k([B, A], [A], 2), 1.0)
|
|
self.assertEqual(recall_at_k([X, Y], [A], 5), 0.0)
|
|
self.assertEqual(recall_at_k([A, B], [A, C], 1), 0.5)
|
|
|
|
def test_ndcg_at_k(self) -> None:
|
|
self.assertEqual(ndcg_at_k([A], [A], 1), 1.0)
|
|
# expected at rank 2: dcg = 1/log2(3), idcg = 1
|
|
self.assertAlmostEqual(
|
|
ndcg_at_k([B, A], [A], 2),
|
|
(1.0 / math.log2(3)) / 1.0,
|
|
)
|
|
self.assertEqual(ndcg_at_k([X, Y], [A], 5), 0.0)
|
|
|
|
|
|
class PermutationTests(unittest.TestCase):
|
|
def test_identical_lists_not_significant(self) -> None:
|
|
scores = [0.5, 1.0, 0.0, 1.0, 0.5]
|
|
report = paired_permutation(scores, list(scores), n_resamples=2000, seed=0)
|
|
self.assertEqual(report["Diff(A-B)"], 0.0)
|
|
self.assertEqual(report["p_value"], 1.0)
|
|
self.assertFalse(report["significant"])
|
|
|
|
def test_large_paired_difference_is_significant(self) -> None:
|
|
a = [1.0] * 20
|
|
b = [0.0] * 20
|
|
report = paired_permutation(a, b, n_resamples=5000, seed=0)
|
|
self.assertGreater(report["Diff(A-B)"], 0.9)
|
|
self.assertLess(report["p_value"], 0.05)
|
|
self.assertTrue(report["significant"])
|
|
|
|
def test_compare_joins_on_retriever_and_query(self) -> None:
|
|
rows_a = (
|
|
[{"retriever": "dense", "query": f"q{i}", "rr": 1.0, "p_at_1": 1} for i in range(20)]
|
|
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
|
|
)
|
|
rows_b = (
|
|
[{"retriever": "dense", "query": f"q{i}", "rr": 0.0, "p_at_1": 0} for i in range(20)]
|
|
+ [{"retriever": "bm25", "query": "q0", "rr": 0.0, "p_at_1": 0}]
|
|
)
|
|
reports = compare(rows_a, rows_b, metric="rr", n_resamples=2000, seed=0)
|
|
by_name = {r["retriever"]: r for r in reports}
|
|
self.assertIn("dense", by_name)
|
|
self.assertTrue(by_name["dense"]["significant"])
|
|
self.assertFalse(by_name["bm25"]["significant"])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|