diff options
Diffstat (limited to 'experiments/resnet_crossover_r2_smoke.py')
| -rw-r--r-- | experiments/resnet_crossover_r2_smoke.py | 118 |
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() |
