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