summaryrefslogtreecommitdiff
path: root/experiments/protocol_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 03:02:55 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 03:02:55 -0500
commit26407ca3e958ef3701cc69afae4973f023b5b19e (patch)
treec9f9d73a0c9c476822fcb2f36b9bc277fee4a670 /experiments/protocol_smoke.py
parent8380cc24a0299b9541b9b619328784ee8ecee206 (diff)
experiments: freeze FashionMNIST C1 recovery
Diffstat (limited to 'experiments/protocol_smoke.py')
-rw-r--r--experiments/protocol_smoke.py12
1 files changed, 12 insertions, 0 deletions
diff --git a/experiments/protocol_smoke.py b/experiments/protocol_smoke.py
index 01c914d..339ac47 100644
--- a/experiments/protocol_smoke.py
+++ b/experiments/protocol_smoke.py
@@ -33,6 +33,17 @@ def main():
assert sum(y.numel() for _, y in validation) == 1000
assert sum(y.numel() for _, y in test) == 10000
+ fmnist = get_dataset_splits(
+ "fmnist", batch_size=256, device="cpu", val_examples=5000, split_seed=2027)
+ f_train, f_validation, f_test, f_n_in, f_n_out, f_metadata = fmnist
+ assert (f_n_in, f_n_out) == (784, 10)
+ assert f_metadata["train_examples"] == 55000
+ assert f_metadata["validation_examples"] == 5000
+ assert f_metadata["validation_class_counts"] == {str(i): 500 for i in range(10)}
+ assert sum(y.numel() for _, y in f_train) == 55000
+ assert sum(y.numel() for _, y in f_validation) == 5000
+ assert sum(y.numel() for _, y in f_test) == 10000
+
probe_loader = get_dataset_splits(
"mnist", batch_size=256, device="cpu", val_examples=1000,
split_seed=2027)[0]
@@ -56,6 +67,7 @@ def main():
assert abs(layerwise["forward_equivalent_batches"] - 11.0 / 3.0) < 1e-12
print("validation split hash:", metadata["validation_index_sha256"])
print("train/validation/test: 59000/1000/10000; stratification exact")
+ print("FashionMNIST recovery split: 55000/5000/10000; stratification exact")
print("training-prefix diagnostic probe: deterministic; shuffle state unchanged")
print("simultaneous/layerwise calibration cost accounting: exact")
print("ALL PROTOCOL SMOKE CHECKS PASSED")