summaryrefslogtreecommitdiff
path: root/experiments/kp_raw_traffic_scaling_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/kp_raw_traffic_scaling_smoke.py')
-rw-r--r--experiments/kp_raw_traffic_scaling_smoke.py52
1 files changed, 52 insertions, 0 deletions
diff --git a/experiments/kp_raw_traffic_scaling_smoke.py b/experiments/kp_raw_traffic_scaling_smoke.py
new file mode 100644
index 0000000..77a5cf6
--- /dev/null
+++ b/experiments/kp_raw_traffic_scaling_smoke.py
@@ -0,0 +1,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()