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()
|