File size: 4,239 Bytes
c87881a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Retrieval metrics with exact expectation over tied-score orderings."""

from __future__ import annotations

import math
from collections import defaultdict
from collections.abc import Iterable, Mapping, Sequence

import numpy as np


def _score_groups(scores: Mapping[str, float]) -> list[tuple[float, list[str]]]:
    grouped: dict[float, list[str]] = defaultdict(list)
    for identifier, score in scores.items():
        if not math.isfinite(float(score)):
            raise ValueError(f"Non-finite retrieval score for {identifier}")
        grouped[float(score)].append(str(identifier))
    return sorted(grouped.items(), key=lambda item: item[0], reverse=True)


def expected_tie_aware_metrics(
    scores: Mapping[str, float],
    relevant_ids: Iterable[str],
    recall_at: Sequence[int] = (10, 50, 100),
    ndcg_at: Sequence[int] = (10, 50, 100),
) -> dict[str, float]:
    relevant = set(relevant_ids)
    missing = relevant.difference(scores)
    if missing:
        raise ValueError(f"Relevant IDs absent from candidates: {sorted(missing)[:10]}")
    if not relevant:
        raise ValueError("At least one relevant candidate is required")

    groups = _score_groups(scores)
    total_candidates = len(scores)
    total_relevant = len(relevant)
    metrics: dict[str, float] = {}

    for cutoff in recall_at:
        effective = min(int(cutoff), total_candidates)
        expected_hits = 0.0
        consumed = 0
        for _, identifiers in groups:
            size = len(identifiers)
            relevant_count = sum(identifier in relevant for identifier in identifiers)
            slots = max(0, min(size, effective - consumed))
            expected_hits += relevant_count * slots / size
            consumed += size
            if consumed >= effective:
                break
        metrics[f"recall@{cutoff}"] = expected_hits / total_relevant
        metrics[f"precision@{cutoff}"] = expected_hits / effective if effective else 0.0

    reciprocal_sum = 0.0
    average_precision_sum = 0.0
    prior_items = 0
    prior_relevant = 0
    for _, identifiers in groups:
        size = len(identifiers)
        relevant_count = sum(identifier in relevant for identifier in identifiers)
        if relevant_count:
            reciprocal_mean = np.mean(
                [1.0 / (prior_items + position) for position in range(1, size + 1)]
            )
            reciprocal_sum += relevant_count * reciprocal_mean
            for position in range(1, size + 1):
                relevant_before = (
                    (position - 1) * (relevant_count - 1) / (size - 1) if size > 1 else 0.0
                )
                probability_relevant = relevant_count / size
                expected_precision = (
                    prior_relevant + 1 + relevant_before
                ) / (prior_items + position)
                average_precision_sum += probability_relevant * expected_precision
        prior_items += size
        prior_relevant += relevant_count
    metrics["mrr"] = reciprocal_sum / total_relevant
    metrics["map"] = average_precision_sum / total_relevant

    for cutoff in ndcg_at:
        effective = min(int(cutoff), total_candidates)
        dcg = 0.0
        consumed = 0
        for _, identifiers in groups:
            size = len(identifiers)
            relevant_count = sum(identifier in relevant for identifier in identifiers)
            slots = max(0, min(size, effective - consumed))
            relevance_probability = relevant_count / size
            for offset in range(slots):
                rank = consumed + offset + 1
                dcg += relevance_probability / math.log2(rank + 1)
            consumed += size
            if consumed >= effective:
                break
        ideal_count = min(effective, total_relevant)
        ideal = sum(1.0 / math.log2(rank + 1) for rank in range(1, ideal_count + 1))
        metrics[f"ndcg@{cutoff}"] = dcg / ideal if ideal else 0.0

    tied_candidates = sum(len(ids) for _, ids in groups if len(ids) > 1)
    metrics["tie_fraction"] = tied_candidates / total_candidates if total_candidates else 0.0
    metrics["candidate_count"] = float(total_candidates)
    metrics["relevant_count"] = float(total_relevant)
    return metrics