diff options
Diffstat (limited to 'sdil/data.py')
| -rw-r--r-- | sdil/data.py | 23 |
1 files changed, 23 insertions, 0 deletions
diff --git a/sdil/data.py b/sdil/data.py index 8dac4df..ad51f59 100644 --- a/sdil/data.py +++ b/sdil/data.py @@ -114,6 +114,29 @@ def _stratified_validation_indices(y, n_val, seed): return train_idx, val_idx +def split_training_loader(loader, val_examples, split_seed, train_batch_size): + """Split an in-memory synthetic training loader without touching its test set.""" + train_idx, val_idx = _stratified_validation_indices( + loader.y.detach().cpu(), val_examples, split_seed) + train_device_idx = train_idx.to(loader.x.device) + val_device_idx = val_idx.to(loader.x.device) + train = _FastLoader(loader.x[train_device_idx], loader.y[train_device_idx], + train_batch_size, True) + validation = _FastLoader(loader.x[val_device_idx], loader.y[val_device_idx], + 1000, False) + digest = hashlib.sha256(val_idx.numpy().tobytes()).hexdigest() + counts = {str(int(cls)): int((loader.y.detach().cpu()[val_idx] == cls).sum()) + for cls in torch.unique(loader.y.detach().cpu()[val_idx], sorted=True)} + return train, validation, { + "split_seed": split_seed, + "validation_examples": int(val_examples), + "validation_index_sha256": digest, + "validation_class_counts": counts, + "train_examples": int(train_idx.numel()), + "split_from_training_only": True, + } + + def get_dataset_splits(name="mnist", batch_size=128, data_dir=DATA_DIR, device="cpu", shuffle_train=True, train_limit=None, val_examples=0, split_seed=2027): """Load real data with an optional training-only validation split. |
