summaryrefslogtreecommitdiff
path: root/experiments/kp_dynamic_projection_confirmation.py
blob: d0bf9c80cac81b14510659e177a06464c288a87b (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
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
#!/usr/bin/env python3
"""Run the frozen D4 paired clean-KP/dynamic-projection test panel."""
import argparse
import json
import os
import subprocess
import sys


CONDITIONS = ("clean_kp", "dynamic")
SEEDS = tuple(range(10, 15))


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "--full_gate", default="results/kp_dynamic_projection_full_gate.json")
    parser.add_argument("--condition", choices=("all",) + CONDITIONS,
                        default="all")
    parser.add_argument("--seed", type=int, choices=SEEDS)
    parser.add_argument("--device", default="cuda")
    parser.add_argument("--dry_run", action="store_true")
    args = parser.parse_args()
    with open(args.full_gate) as handle:
        gate = json.load(handle)
    if (gate.get("protocol") != "kp_dynamic_neutral_projection_full_v1"
            or gate.get("status") != "passed"
            or gate.get("independent_confirmation_opened") is not True):
        raise ValueError("D4 requires the audited D3 pass")

    conditions = CONDITIONS if args.condition == "all" else (args.condition,)
    seeds = SEEDS if args.seed is None else (args.seed,)
    os.makedirs("results/kp_dynamic_projection_confirmation", exist_ok=True)
    for seed in seeds:
        for condition in conditions:
            dynamic = condition == "dynamic"
            command = [
                sys.executable, "experiments/conv_run.py",
                "--mode", "kp_traffic" if dynamic else "kp",
            ]
            if dynamic:
                command.extend([
                    "--traffic_rule", "innovation",
                    "--predictor_mode", "closed_form",
                    "--neutral_projection", "1",
                    "--traffic_seed", str(5000 + seed),
                    "--traffic_ratio", "4",
                    "--traffic_calibration_examples", "64",
                    "--learn_P", "1", "--eta_P", "0.1",
                    "--predictor_warmup_steps", "1",
                    "--predictor_every", "0",
                ])
            command.extend([
                "--device", args.device, "--depth", "20", "--width", "16",
                "--seed", str(seed), "--loader_seed", str(seed),
                "--batch_size", "128", "--epochs", "200",
                "--train_limit", "0", "--val_examples", "0",
                "--split_seed", "2027", "--eval_split", "test",
                "--eval_every", "0", "--augment_train", "1",
                "--lr", "0.1", "--output_lr", "0.1",
                "--lr_schedule", "step", "--lr_milestones", "100,150",
                "--lr_gamma", "0.1", "--warmup_epochs", "0",
                "--momentum", "0.9", "--weight_decay", "1e-4",
                "--normalization", "batchnorm", "--a_scale", "1",
                "--alignment_probe", "32",
                "--out", ("results/kp_dynamic_projection_confirmation/"
                          f"seed{seed}_{condition}.json"),
            ])
            print(" ".join(command), flush=True)
            if not args.dry_run:
                subprocess.run(command, check=True)


if __name__ == "__main__":
    main()