diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:45:36 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:45:36 -0500 |
| commit | c8dd2e591c9835b8618fc34bc1d043634ce345e1 (patch) | |
| tree | 3f1ec546b4832eab5147cb2baaa92605f602b24d /experiments/diagnose_kp_traffic_nonfinite.py | |
| parent | 72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (diff) | |
experiment: freeze predictor after neutral fit
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 7 |
1 files changed, 6 insertions, 1 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py index bdf79c3..99a4e99 100644 --- a/experiments/diagnose_kp_traffic_nonfinite.py +++ b/experiments/diagnose_kp_traffic_nonfinite.py @@ -89,12 +89,15 @@ def main(): parser.add_argument("--predictor_mode", choices=("nlms20", "closed_form"), default="nlms20") + parser.add_argument("--predictor_every", type=int, default=16) parser.add_argument("--device", default="cuda") parser.add_argument("--max_steps", type=int, default=352) parser.add_argument("--out", required=True) args = parser.parse_args() if args.max_steps < 1: raise ValueError("max_steps must be positive") + if args.predictor_every < 0: + raise ValueError("predictor_every must be nonnegative") torch.manual_seed(0) if str(args.device).startswith("cuda"): @@ -168,7 +171,8 @@ def main(): if step >= args.max_steps: break result = conv_kp_mixed_traffic_step( - net, x, y, config, step, args.rule, predictor_every=16) + net, x, y, config, step, args.rule, + predictor_every=args.predictor_every) state = network_state(net) bad = nonfinite_groups(state) for name, indices in bad.items(): @@ -223,6 +227,7 @@ def main(): "provenance": provenance(), "rule": args.rule, "predictor_mode": args.predictor_mode, + "predictor_every": args.predictor_every, "max_steps": args.max_steps, "split": split, "traffic_calibration": calibration, |
