whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
3.79 kB
"""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