Spaces:
Sleeping
Sleeping
| """Split validation helpers.""" | |
| from __future__ import annotations | |
| import pandas as pd | |
| def split_cat_sets(split_df: pd.DataFrame) -> tuple[set[str], set[str]]: | |
| """Return train/test cat sets for the cat-heldout split.""" | |
| train_cats = set(split_df.loc[split_df["cat_heldout_split"] == "train", "cat_id"]) | |
| test_cats = set(split_df.loc[split_df["cat_heldout_split"] == "test", "cat_id"]) | |
| return train_cats, test_cats | |
| def cat_overlap(split_df: pd.DataFrame) -> set[str]: | |
| """Return cats appearing in both cat-heldout train and test.""" | |
| train_cats, test_cats = split_cat_sets(split_df) | |
| return train_cats & test_cats | |
| def assert_zero_cat_overlap(split_df: pd.DataFrame) -> None: | |
| """Raise when the cat-heldout split is not subject-disjoint.""" | |
| overlap = cat_overlap(split_df) | |
| if overlap: | |
| raise AssertionError(f"expected zero cat overlap, found {sorted(overlap)}") | |