diff options
| -rw-r--r-- | experiments/conv_local_smoke.py | 11 | ||||
| -rw-r--r-- | experiments/diagnose_kp_traffic_nonfinite.py | 7 | ||||
| -rw-r--r-- | sdil/conv.py | 7 |
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( |
