whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
4.24 kB
"""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