From 72f6b758540e8c1d7a44fed63bafe71784f6dfa2 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 16:43:39 -0500 Subject: experiment: add closed-form neutral innovation fit --- sdil/conv.py | 55 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 55 insertions(+) (limited to 'sdil/conv.py') diff --git a/sdil/conv.py b/sdil/conv.py index ea62e95..d6bae4f 100644 --- a/sdil/conv.py +++ b/sdil/conv.py @@ -768,6 +768,61 @@ class CIFARKPMixedTrafficResNet(CIFARKPResNet): units += residual.numel() return squared_error / units + @torch.no_grad() + def predictor_closed_form_fit(self, hiddens, min_variance=1e-12): + """Fit the local affine neutral relation by per-cell least squares. + + Each spatial cell uses only its own soma and instruction-off apical + observations across the supplied batch. Cells with no observed soma + variance receive a zero slope and their local target mean as bias. + """ + if min_variance < 0: + raise ValueError("minimum predictor variance must be nonnegative") + traffic = self.traffic_fields(hiddens) + residual_power = 0.0 + traffic_power = 0.0 + maximum_residual_slope = 0.0 + squared_error = 0.0 + units = 0 + for hidden, target, slope, bias in zip( + hiddens, traffic, self.P_traffic, self.P_traffic_bias): + hidden_mean = hidden.mean(dim=0) + target_mean = target.mean(dim=0) + centered_h = hidden - hidden_mean + centered_target = target - target_mean + variance = centered_h.square().mean(dim=0) + covariance = (centered_h * centered_target).mean(dim=0) + active = variance > min_variance + 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) + bias.copy_(fitted_bias) + residual = target - slope * hidden - bias + residual_covariance = (centered_h * ( + residual - residual.mean(dim=0))).mean(dim=0) + residual_slope = torch.where( + active, + residual_covariance / variance.clamp_min(min_variance), + torch.zeros_like(variance)) + maximum_residual_slope = max( + maximum_residual_slope, float(residual_slope.abs().max())) + power = float(residual.square().sum()) + residual_power += power + traffic_power += float(target.square().sum()) + squared_error += power + units += residual.numel() + if traffic_power <= 0: + raise ValueError("closed-form predictor fit requires nonzero traffic") + return { + "mse": squared_error / units, + "residual_traffic_rms_ratio": math.sqrt( + residual_power / traffic_power), + "max_absolute_residual_soma_slope": maximum_residual_slope, + "observations": int(hiddens[0].shape[0]), + } + @torch.no_grad() def predictor_traffic_residual_rms_ratio(self, hiddens): traffic = self.traffic_fields(hiddens) -- cgit v1.2.3