summaryrefslogtreecommitdiff
path: root/experiments/kp_raw_traffic_scaling_smoke.py
blob: 77a5cf6881c3494c4f4c5b42ef910722b56a77f5 (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
#!/usr/bin/env python3
"""Registry smoke for the frozen raw-KP scaling control."""
import os
import sys


sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from experiments.kp_raw_traffic_scaling import DEPTHS, jobs


def value(command, flag):
    return command[command.index(flag) + 1]


def main():
    records = jobs()
    assert len(records) == 3
    assert {record["depth"] for record in records} == set(DEPTHS)
    assert len({record["output"] for record in records}) == 3
    for record in records:
        command = record["command"]
        expected = {
            "--mode": "kp_traffic",
            "--traffic_rule": "raw",
            "--traffic_ratio": "4",
            "--traffic_seed": "5000",
            "--traffic_calibration_examples": "64",
            "--predictor_mode": "closed_form",
            "--predictor_warmup_steps": "1",
            "--predictor_every": "0",
            "--neutral_projection": "1",
            "--depth": str(record["depth"]),
            "--width": "16",
            "--seed": "0",
            "--loader_seed": "0",
            "--split_seed": "2027",
            "--epochs": "200",
            "--eval_split": "validation",
            "--eval_every": "1",
            "--lr": "0.1",
            "--output_lr": "0.1",
            "--lr_schedule": "step",
            "--lr_milestones": "100,150",
        }
        for flag, wanted in expected.items():
            assert value(command, flag) == wanted, (flag, command)
        assert record["timeout_seconds"] == 48 * 60 * 60
    print("raw-KP traffic scaling registry: 3/3 exact")


if __name__ == "__main__":
    main()