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