summaryrefslogtreecommitdiff
path: root/experiments/resnet_crossover_r2_smoke.py
blob: d98c638db9837a6d4ebf85de747e5a9f6143c2d9 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
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()