summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
Diffstat (limited to 'sdil')
-rw-r--r--sdil/data.py23
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.