"""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