File size: 3,794 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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""Deterministic group-level split construction and leakage assertions."""

from __future__ import annotations

import hashlib
import json
import random
from pathlib import Path

import pandas as pd

from .artifacts import sha256_file
from .data import require_columns


SPLITS = ("train", "validation", "test")


def build_group_split(
    labels: pd.DataFrame,
    seed: int,
    minimum_group_size: int = 3,
    fractions: tuple[float, float, float] = (0.70, 0.15, 0.15),
) -> pd.DataFrame:
    require_columns(labels, ["bgc_id", "group_id", "label_tier"], "labels")
    if abs(sum(fractions) - 1.0) > 1e-9 or any(value <= 0 for value in fractions):
        raise ValueError("Split fractions must be positive and sum to one")
    if labels["bgc_id"].duplicated().any():
        raise ValueError("Each BGC must have exactly one group assignment")

    sizes = labels.groupby("group_id")["bgc_id"].nunique()
    eligible = sorted(sizes[sizes >= minimum_group_size].index.astype(str))
    if len(eligible) < 3:
        raise ValueError("At least three eligible groups are required")
    random.Random(seed).shuffle(eligible)

    train_end = max(1, round(len(eligible) * fractions[0]))
    validation_end = max(train_end + 1, round(len(eligible) * sum(fractions[:2])))
    validation_end = min(validation_end, len(eligible) - 1)
    split_groups = {
        "train": set(eligible[:train_end]),
        "validation": set(eligible[train_end:validation_end]),
        "test": set(eligible[validation_end:]),
    }
    group_to_split = {
        group: split_name for split_name, groups in split_groups.items() for group in groups
    }
    result = labels[labels["group_id"].astype(str).isin(group_to_split)].copy()
    result["group_id"] = result["group_id"].astype(str)
    result["split"] = result["group_id"].map(group_to_split)
    result = result.sort_values(["split", "group_id", "bgc_id"]).reset_index(drop=True)
    validate_split(result)
    return result


def validate_split(assignments: pd.DataFrame) -> None:
    require_columns(assignments, ["bgc_id", "group_id", "split", "label_tier"], "split")
    invalid = set(assignments["split"]).difference(SPLITS)
    if invalid:
        raise ValueError(f"Unknown split names: {sorted(invalid)}")
    if assignments["bgc_id"].duplicated().any():
        raise ValueError("BGC identifiers overlap within the split manifest")
    memberships = assignments.groupby("group_id")["split"].nunique()
    leaking = memberships[memberships > 1]
    if not leaking.empty:
        raise ValueError(f"Groups cross split boundaries: {leaking.index.tolist()[:10]}")
    if set(assignments["split"]) != set(SPLITS):
        raise ValueError("Train, validation, and test must all be non-empty")


def split_summary(assignments: pd.DataFrame) -> dict[str, object]:
    validate_split(assignments)
    by_split = {}
    for split_name, frame in assignments.groupby("split"):
        by_split[str(split_name)] = {
            "bgcs": int(frame["bgc_id"].nunique()),
            "groups": int(frame["group_id"].nunique()),
            "label_tiers": frame["label_tier"].value_counts().to_dict(),
        }
    serial = assignments.sort_values("bgc_id").to_dict("records")
    fingerprint = hashlib.sha256(
        json.dumps(serial, sort_keys=True, separators=(",", ":")).encode("utf-8")
    ).hexdigest()
    return {"schema_version": 1, "split_sha256": fingerprint, "by_split": by_split}


def load_split(path: str | Path, expected_sha256: str | None = None) -> pd.DataFrame:
    split_path = Path(path)
    if expected_sha256 and sha256_file(split_path) != expected_sha256:
        raise ValueError("Split file fingerprint does not match the expected value")
    assignments = pd.read_csv(split_path)
    validate_split(assignments)
    return assignments