summaryrefslogtreecommitdiff
path: root/experiments/resnet_crossover_r2_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/resnet_crossover_r2_smoke.py')
-rw-r--r--experiments/resnet_crossover_r2_smoke.py118
1 files changed, 118 insertions, 0 deletions
diff --git a/experiments/resnet_crossover_r2_smoke.py b/experiments/resnet_crossover_r2_smoke.py
new file mode 100644
index 0000000..d98c638
--- /dev/null
+++ b/experiments/resnet_crossover_r2_smoke.py
@@ -0,0 +1,118 @@
+#!/usr/bin/env python3
+"""Deterministic registry audit for the complete ResNet R2 crossover."""
+import os
+import sys
+from types import SimpleNamespace
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from experiments.resnet_crossover_r2 import DEPTHS, METHODS, r2_jobs
+from experiments.resnet_crossover_native import work_report
+
+
+RATES = {
+ "bp": 0.1,
+ "fa": 0.03,
+ "dfa": 0.03,
+ "pepita": 0.001,
+ "ff": 0.03,
+ "ep": 0.001,
+ "dualprop": 0.025,
+ "clean_kp": 0.1,
+ "sdil": 0.1,
+}
+
+
+def value(command, flag):
+ index = command.index(flag)
+ return command[index + 1]
+
+
+def main():
+ jobs = r2_jobs(RATES)
+ assert len(jobs) == 27
+ assert {(job["method"], job["depth"]) for job in jobs} == {
+ (method, depth) for method in METHODS for depth in DEPTHS
+ }
+ assert len({job["experiment_name"] for job in jobs}) == 27
+ assert len({job["output"] for job in jobs}) == 27
+ for job in jobs:
+ method = job["method"]
+ command = job["command"]
+ expected_epochs = (
+ 100 if method in ("pepita", "ep")
+ else 40 if method == "ff"
+ else 200
+ )
+ expected_schedule = (
+ "pepita" if method == "pepita"
+ else "constant" if method in ("ff", "ep")
+ else "step"
+ )
+ assert job["epochs"] == expected_epochs
+ assert job["lr_schedule"] == expected_schedule
+ assert value(command, "--depth") == str(job["depth"])
+ assert value(command, "--epochs") == str(expected_epochs)
+ assert value(command, "--lr_schedule") == expected_schedule
+ assert value(command, "--eval_split") == "validation"
+ assert value(command, "--val_examples") == "5000"
+ assert value(command, "--eval_every") == "1"
+ assert value(command, "--lr") == str(RATES[method])
+ assert "--test" not in command
+ assert job["timeout_seconds"] == 48 * 60 * 60
+ if method in ("fa", "dfa"):
+ assert value(command, "--output_lr") == "0.1"
+ else:
+ assert value(command, "--output_lr") == str(RATES[method])
+ if method == "pepita":
+ assert value(command, "--lr_milestones") == "60,90"
+ elif method not in ("ff", "ep"):
+ assert value(command, "--lr_milestones") == "100,150"
+ if method == "sdil":
+ assert value(command, "--traffic_ratio") == "4"
+ assert value(command, "--neutral_projection") == "1"
+ assert value(command, "--alignment_probe") == "32"
+ if method == "ep":
+ assert value(command, "--ep_free_steps") == "20"
+ assert value(command, "--ep_nudge_steps") == "4"
+ if method == "dualprop":
+ assert value(command, "--dp_inference_passes") == "16"
+ shard0 = jobs[0::2]
+ shard1 = jobs[1::2]
+ assert len(shard0) == 14 and len(shard1) == 13
+ assert any(job["method"] == "ep" for job in shard0)
+ assert any(job["method"] == "ep" for job in shard1)
+ assert any(job["method"] == "dualprop" for job in shard0)
+ assert any(job["method"] == "dualprop" for job in shard1)
+ assert any(job["method"] == "ff" for job in shard0)
+ assert any(job["method"] == "ff" for job in shard1)
+ dummy_net = SimpleNamespace(
+ n_hidden=19,
+ n_forward_parameters=123,
+ forward_macs_per_example=456,
+ )
+ for method in ("pepita", "ff", "ep", "dualprop"):
+ args = SimpleNamespace(
+ method=method,
+ ep_free_steps=20,
+ ep_nudge_steps=4,
+ dp_inference_passes=16,
+ )
+ work = work_report(dummy_net, args, 100, 20, 1)
+ assert work["logical_task_loss_queries"] == 0
+ if method == "ff":
+ assert work["local_target_backward_example_evaluations"] == 200
+ assert work["candidate_label_evaluation_presentations"] == 200
+ elif method == "ep":
+ assert work["relaxation_example_passes"] == 2800
+ assert work["local_vjp_example_evaluations"] == 2800
+ elif method == "dualprop":
+ assert work["relaxation_example_passes"] == 1600
+ assert work["local_vjp_example_evaluations"] == 1600
+ else:
+ assert work["training_example_presentations"] == 200
+ assert work["local_vjp_example_evaluations"] == 0
+ print("ResNet R2 registry: 27/27 cells; schedules and shards exact")
+
+
+if __name__ == "__main__":
+ main()