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