| """Locked split retrieval protocol shared by learned and baseline embeddings.""" |
|
|
| from __future__ import annotations |
|
|
| import random |
| from collections import defaultdict |
| from collections.abc import Callable, Mapping, Sequence |
|
|
| import pandas as pd |
|
|
| from .metrics import expected_tie_aware_metrics |
| from .splits import validate_split |
|
|
|
|
| ScoreFunction = Callable[[Sequence[str], Sequence[str]], Mapping[str, float]] |
|
|
|
|
| def evaluate_retrieval( |
| assignments: pd.DataFrame, |
| split_name: str, |
| methods: Mapping[str, ScoreFunction], |
| reference_size: int, |
| draws: int, |
| seed: int, |
| recall_at: Sequence[int] = (10, 50, 100), |
| ndcg_at: Sequence[int] = (10, 50, 100), |
| ) -> pd.DataFrame: |
| validate_split(assignments) |
| split = assignments[assignments["split"] == split_name] |
| universe = sorted(split["bgc_id"].astype(str).unique()) |
| group_members: dict[str, list[str]] = defaultdict(list) |
| for row in split.itertuples(index=False): |
| group_members[str(row.group_id)].append(str(row.bgc_id)) |
| eligible = { |
| group: sorted(set(members)) |
| for group, members in group_members.items() |
| if len(set(members)) > reference_size |
| } |
| if not eligible: |
| raise ValueError(f"No {split_name} groups have more than {reference_size} members") |
|
|
| rows: list[dict[str, object]] = [] |
| for group_id, members in sorted(eligible.items()): |
| for draw in range(draws): |
| draw_seed = f"{seed}:{split_name}:{group_id}:{draw}" |
| references = random.Random(draw_seed).sample(members, reference_size) |
| relevant = set(members).difference(references) |
| candidates = [identifier for identifier in universe if identifier not in references] |
| for method_name, score_function in methods.items(): |
| scores = dict(score_function(candidates, references)) |
| if set(scores) != set(candidates): |
| raise ValueError(f"{method_name} did not score the complete candidate universe") |
| metrics = expected_tie_aware_metrics(scores, relevant, recall_at, ndcg_at) |
| rows.append( |
| { |
| "split": split_name, |
| "group_id": group_id, |
| "draw": draw, |
| "method": method_name, |
| "reference_ids": ";".join(references), |
| **metrics, |
| } |
| ) |
| return pd.DataFrame(rows) |
|
|