summaryrefslogtreecommitdiff
path: root/experiments/diagnose_kp_traffic_nonfinite.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:43:39 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:43:39 -0500
commit72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (patch)
treede271bb7948e31972e76070518677c12c4f79058 /experiments/diagnose_kp_traffic_nonfinite.py
parent575530e7124c540e58571b50ef147a5f20e6cb54 (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.py30
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,