diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:43:39 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 16:43:39 -0500 |
| commit | 72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (patch) | |
| tree | de271bb7948e31972e76070518677c12c4f79058 /experiments/diagnose_kp_traffic_nonfinite.py | |
| parent | 575530e7124c540e58571b50ef147a5f20e6cb54 (diff) | |
experiment: add closed-form neutral innovation fit
Diffstat (limited to 'experiments/diagnose_kp_traffic_nonfinite.py')
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 30 |
1 files changed, 21 insertions, 9 deletions
diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py index a045a75..bdf79c3 100644 --- a/experiments/diagnose_kp_traffic_nonfinite.py +++ b/experiments/diagnose_kp_traffic_nonfinite.py @@ -86,6 +86,9 @@ def main(): parser = argparse.ArgumentParser() parser.add_argument("--rule", choices=("raw", "matched", "innovation"), default="innovation") + parser.add_argument("--predictor_mode", + choices=("nlms20", "closed_form"), + default="nlms20") parser.add_argument("--device", default="cuda") parser.add_argument("--max_steps", type=int, default=352) parser.add_argument("--out", required=True) @@ -127,21 +130,28 @@ def main(): calibration_error, calibration_forward) calibration = net.calibrate_traffic_gain( calibration_instruction, calibration_forward["hiddens"], 4.0) - del calibration_forward, calibration_error, calibration_instruction - - iterator = iter(train) + closed_form_fit = None warmup_mse = [] - for _ in range(20): - x, _ = next(iterator) - forward = net.forward(x, training=True, update_stats=False) - warmup_mse.append(net.predictor_step(forward["hiddens"], 0.1)) + if args.predictor_mode == "closed_form": + closed_form_fit = net.predictor_closed_form_fit( + calibration_forward["hiddens"]) + warmup_mse.append(closed_form_fit["mse"]) + else: + iterator = iter(train) + for _ in range(20): + x, _ = next(iterator) + forward = net.forward(x, training=True, update_stats=False) + warmup_mse.append(net.predictor_step(forward["hiddens"], 0.1)) + del calibration_forward, calibration_error, calibration_instruction train.g.set_state(loader_state) loader_state_restored = torch.equal(train.g.get_state(), loader_state) audit_forward = net.forward( calibration_x, training=True, update_stats=False) post_warmup_ratio = net.predictor_traffic_residual_rms_ratio( audit_forward["hiddens"]) - del forward, audit_forward + if args.predictor_mode == "nlms20": + del forward + del audit_forward initial_state = network_state(net) if nonfinite_groups(initial_state): @@ -212,15 +222,17 @@ def main(): "scope": "training_only_no_validation_or_test_evaluation", "provenance": provenance(), "rule": args.rule, + "predictor_mode": args.predictor_mode, "max_steps": args.max_steps, "split": split, "traffic_calibration": calibration, "predictor_warmup": { - "steps": 20, + "steps": len(warmup_mse), "first_mse": warmup_mse[0], "last_mse": warmup_mse[-1], "post_warmup_traffic_residual_rms_ratio": post_warmup_ratio, "task_loader_state_restored": loader_state_restored, + "closed_form_fit": closed_form_fit, }, "initial_parameter_state": initial_state, "trajectory": trajectory, |
