File size: 2,491 Bytes
c87881a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
"""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)