from sklearn.model_selection import GroupShuffleSplit def assign_splits(manifest, val_fraction=0.15, seed=42): train_data = manifest[manifest["orig_split"] == "train"] groups = train_data["subject_id"].values gss = GroupShuffleSplit(n_splits=1, test_size=val_fraction, random_state=seed) train_idx, val_idx = next(gss.split(X=train_data, y=None, groups=groups)) train_subjects = set(train_data.iloc[train_idx]["subject_id"].unique()) val_subjects = set(train_data.iloc[val_idx]["subject_id"].unique()) # Crash loudly if leakage ever sneaks in assert train_subjects.isdisjoint(val_subjects), "Subject leak detected!" return train_subjects, val_subjects