diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-27 10:35:06 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-27 10:35:06 -0500 |
| commit | 2b9caf80d9a1054ec585a03718fc37106f8c7aae (patch) | |
| tree | 230c9d478498fbce5249edbe74b1abbccdabbfe4 /experiments/kp_raw_traffic_scaling_smoke.py | |
| parent | db21295addcfeb830ae72b330a8e7784c75c5232 (diff) | |
exp: freeze raw KP traffic scaling control
Diffstat (limited to 'experiments/kp_raw_traffic_scaling_smoke.py')
| -rw-r--r-- | experiments/kp_raw_traffic_scaling_smoke.py | 52 |
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() |
