#!/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()