summaryrefslogtreecommitdiff
path: root/sdil/data.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:47:20 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:47:20 -0500
commit296c95045822c6d13de35515728528e411a7f8cd (patch)
treebe02a1c0053b1658781393669098ca68ff090a63 /sdil/data.py
parent4f0c65d252d21438da48974ac77deb4a862abcd3 (diff)
experiments: add synthetic validation isolation
Diffstat (limited to 'sdil/data.py')
-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.