summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:45:36 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:45:36 -0500
commitc8dd2e591c9835b8618fc34bc1d043634ce345e1 (patch)
tree3f1ec546b4832eab5147cb2baaa92605f602b24d
parent72f6b758540e8c1d7a44fed63bafe71784f6dfa2 (diff)
experiment: freeze predictor after neutral fit
-rw-r--r--experiments/conv_local_smoke.py11
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py7
-rw-r--r--sdil/conv.py7
3 files changed, 21 insertions, 4 deletions
diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py
index 3267417..802051b 100644
--- a/experiments/conv_local_smoke.py
+++ b/experiments/conv_local_smoke.py
@@ -1075,6 +1075,17 @@ def kp_mixed_traffic_checks():
weight_decay=0.0, learn_A=False, learn_P=True),
step=0, rule="innovation", predictor_every=16)
assert math.isfinite(result["loss"]) and result["did_predictor_update"]
+ frozen_predictor = [value.clone() for value in
+ net.P_traffic + net.P_traffic_bias]
+ frozen_result = conv_kp_mixed_traffic_step(
+ net, x, y, ConvSDILConfig(
+ eta=1e-4, eta_output=1e-4, eta_P=0.1, momentum=0.0,
+ weight_decay=0.0, learn_A=False, learn_P=True),
+ step=1, rule="innovation", predictor_every=0)
+ assert math.isfinite(frozen_result["loss"])
+ assert not frozen_result["did_predictor_update"]
+ assert all(torch.equal(before, after) for before, after in zip(
+ frozen_predictor, net.P_traffic + net.P_traffic_bias))
assert all(not value.requires_grad for value in
net.W + net.Q + net.P_traffic + net.P_traffic_bias
+ [net.W_out, net.R_out, net.b_out])
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,
diff --git a/sdil/conv.py b/sdil/conv.py
index d6bae4f..285ceb1 100644
--- a/sdil/conv.py
+++ b/sdil/conv.py
@@ -1267,8 +1267,8 @@ def conv_kp_mixed_traffic_step(net, x, y, config, step, rule,
"""One reciprocal-KP update using a selected mixed-apical signal."""
if not isinstance(net, CIFARKPMixedTrafficResNet):
raise TypeError("mixed-traffic KP step requires CIFARKPMixedTrafficResNet")
- if predictor_every < 1:
- raise ValueError("predictor cadence must be positive")
+ if predictor_every < 0:
+ raise ValueError("predictor cadence must be nonnegative")
config.validate()
with torch.no_grad():
forward = net.forward(
@@ -1303,7 +1303,8 @@ def conv_kp_mixed_traffic_step(net, x, y, config, step, rule,
momentum=config.momentum, weight_decay=config.weight_decay,
gamma_directions=gamma_directions,
beta_directions=beta_directions)
- did_predictor_update = step % predictor_every == 0
+ did_predictor_update = (
+ predictor_every > 0 and step % predictor_every == 0)
predictor_mse = None
if did_predictor_update:
predictor_mse = net.predictor_step(