bgc-setnet / source /src /bgc_retrieval /evaluation.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
2.49 kB
"""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)