summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:47:55 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 16:47:55 -0500
commit293785af4c9a556f388f04bcf9599aaee9e9dfd9 (patch)
tree1234f2dba19c67d5586a226a13f2e954910623c8
parentc8dd2e591c9835b8618fc34bc1d043634ce345e1 (diff)
experiment: add one-sided residual stability margin
-rw-r--r--experiments/conv_local_smoke.py9
-rw-r--r--experiments/diagnose_kp_traffic_nonfinite.py7
-rw-r--r--sdil/conv.py26
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
@@ -997,6 +997,15 @@ def kp_mixed_traffic_checks():
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_()
components = net.mixed_apical_components(
instruction, forward["hiddens"], "matched")
norm_errors = []
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]),
}