From 293785af4c9a556f388f04bcf9599aaee9e9dfd9 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 16:47:55 -0500 Subject: experiment: add one-sided residual stability margin --- experiments/conv_local_smoke.py | 9 +++++++++ experiments/diagnose_kp_traffic_nonfinite.py | 7 ++++++- sdil/conv.py | 26 +++++++++++++++++++++++--- 3 files changed, 38 insertions(+), 4 deletions(-) diff --git a/experiments/conv_local_smoke.py b/experiments/conv_local_smoke.py index 802051b..dbae759 100644 --- a/experiments/conv_local_smoke.py +++ b/experiments/conv_local_smoke.py @@ -994,6 +994,15 @@ def kp_mixed_traffic_checks(): assert closed_form["residual_traffic_rms_ratio"] < 1e-14 assert closed_form["max_absolute_residual_soma_slope"] < 1e-14 + for slope, bias in zip(net.P_traffic, net.P_traffic_bias): + slope.zero_() + bias.zero_() + stable_fit = net.predictor_closed_form_fit( + forward["hiddens"], stability_margin=1e-3) + assert stable_fit["max_positive_residual_soma_slope"] < 1e-14 + assert stable_fit["min_residual_soma_slope"] < -9e-4 + assert stable_fit["max_applied_stability_margin"] >= 1e-3 + for slope, bias in zip(net.P_traffic, net.P_traffic_bias): slope.zero_() bias.zero_() diff --git a/experiments/diagnose_kp_traffic_nonfinite.py b/experiments/diagnose_kp_traffic_nonfinite.py index 99a4e99..7b7389f 100644 --- a/experiments/diagnose_kp_traffic_nonfinite.py +++ b/experiments/diagnose_kp_traffic_nonfinite.py @@ -90,6 +90,7 @@ def main(): choices=("nlms20", "closed_form"), default="nlms20") parser.add_argument("--predictor_every", type=int, default=16) + parser.add_argument("--stability_margin", type=float, default=0.0) parser.add_argument("--device", default="cuda") parser.add_argument("--max_steps", type=int, default=352) parser.add_argument("--out", required=True) @@ -98,6 +99,8 @@ def main(): raise ValueError("max_steps must be positive") if args.predictor_every < 0: raise ValueError("predictor_every must be nonnegative") + if args.stability_margin < 0: + raise ValueError("stability_margin must be nonnegative") torch.manual_seed(0) if str(args.device).startswith("cuda"): @@ -137,7 +140,8 @@ def main(): warmup_mse = [] if args.predictor_mode == "closed_form": closed_form_fit = net.predictor_closed_form_fit( - calibration_forward["hiddens"]) + calibration_forward["hiddens"], + stability_margin=args.stability_margin) warmup_mse.append(closed_form_fit["mse"]) else: iterator = iter(train) @@ -228,6 +232,7 @@ def main(): "rule": args.rule, "predictor_mode": args.predictor_mode, "predictor_every": args.predictor_every, + "stability_margin": args.stability_margin, "max_steps": args.max_steps, "split": split, "traffic_calibration": calibration, diff --git a/sdil/conv.py b/sdil/conv.py index 285ceb1..d9364a2 100644 --- a/sdil/conv.py +++ b/sdil/conv.py @@ -769,7 +769,8 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): return squared_error / units @torch.no_grad() - def predictor_closed_form_fit(self, hiddens, min_variance=1e-12): + def predictor_closed_form_fit(self, hiddens, min_variance=1e-12, + stability_margin=0.0): """Fit the local affine neutral relation by per-cell least squares. Each spatial cell uses only its own soma and instruction-off apical @@ -778,10 +779,15 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): """ if min_variance < 0: raise ValueError("minimum predictor variance must be nonnegative") + if stability_margin < 0: + raise ValueError("predictor stability margin must be nonnegative") traffic = self.traffic_fields(hiddens) residual_power = 0.0 traffic_power = 0.0 maximum_residual_slope = 0.0 + maximum_positive_residual_slope = 0.0 + minimum_residual_slope = 0.0 + maximum_applied_margin = 0.0 squared_error = 0.0 units = 0 for hidden, target, slope, bias in zip( @@ -796,8 +802,10 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): fitted_slope = torch.where( active, covariance / variance.clamp_min(min_variance), torch.zeros_like(variance)) - fitted_bias = target_mean - fitted_slope * hidden_mean - slope.copy_(fitted_slope) + applied_margin = stability_margin * (1.0 + fitted_slope.abs()) + stabilized_slope = fitted_slope + applied_margin + fitted_bias = target_mean - stabilized_slope * hidden_mean + slope.copy_(stabilized_slope) bias.copy_(fitted_bias) residual = target - slope * hidden - bias residual_covariance = (centered_h * ( @@ -808,6 +816,13 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): torch.zeros_like(variance)) maximum_residual_slope = max( maximum_residual_slope, float(residual_slope.abs().max())) + maximum_positive_residual_slope = max( + maximum_positive_residual_slope, + float(residual_slope.max().clamp_min(0.0))) + minimum_residual_slope = min( + minimum_residual_slope, float(residual_slope.min())) + maximum_applied_margin = max( + maximum_applied_margin, float(applied_margin.max())) power = float(residual.square().sum()) residual_power += power traffic_power += float(target.square().sum()) @@ -820,6 +835,11 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): "residual_traffic_rms_ratio": math.sqrt( residual_power / traffic_power), "max_absolute_residual_soma_slope": maximum_residual_slope, + "max_positive_residual_soma_slope": ( + maximum_positive_residual_slope), + "min_residual_soma_slope": minimum_residual_slope, + "max_applied_stability_margin": maximum_applied_margin, + "stability_margin": float(stability_margin), "observations": int(hiddens[0].shape[0]), } -- cgit v1.2.3