summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:14:42 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:14:42 -0500
commit200625df021c02e83aeecd496b4d4f5f3ffe8ad5 (patch)
tree126decd39e5e6fcedecba89ee698589d8a89b072 /sdil
parentbcd845b60dc85e4f46bf1ab343405c1e40ed0860 (diff)
oral-a: make BatchNorm credit local and causal
Diffstat (limited to 'sdil')
-rw-r--r--sdil/conv.py218
1 files changed, 179 insertions, 39 deletions
diff --git a/sdil/conv.py b/sdil/conv.py
index 996817b..3b8fde1 100644
--- a/sdil/conv.py
+++ b/sdil/conv.py
@@ -1,11 +1,11 @@
"""Convolutional local-learning primitives for CIFAR residual networks.
The forward topology is the standard CIFAR ``6n+2`` basic-block family with
-option-A identity shortcuts. Batch normalization is deliberately absent:
-its cross-example Jacobian obscures what information a synapse needs. A
-fixed ``1/sqrt(number of blocks)`` residual multiplier keeps the otherwise
-normalization-free network stable and is included explicitly in every local
-eligibility calculation.
+option-A identity shortcuts. It supports canonical BatchNorm as well as a
+normalization-free ablation. BatchNorm's cross-example Jacobian is evaluated
+inside the current layer only; the synaptic update still never reads a
+downstream weight. Normalization-free networks use an explicit residual
+multiplier, which is included in every local eligibility calculation.
Forward parameters are plain tensors. The local rule uses only the stored
presynaptic activation, a postsynaptic ReLU gate, and a teaching vector at the
@@ -33,7 +33,7 @@ class ConvLayerSpec:
class CIFARLocalResNet:
- """Normalization-free CIFAR ResNet with explicit local eligibilities.
+ """CIFAR ResNet with explicit local eligibilities.
``depth`` must satisfy ``depth = 6n + 2``. Hidden populations are defined
after the stem ReLU, after every block's first ReLU, and after every block
@@ -43,7 +43,8 @@ class CIFARLocalResNet:
def __init__(self, depth=20, base_width=16, n_classes=10, device="cpu",
dtype=torch.float32, seed=0, weight_scale=1.0,
- residual_scale=None):
+ residual_scale=None, normalization="none", bn_momentum=0.1,
+ bn_eps=1e-5):
if depth < 8 or (depth - 2) % 6:
raise ValueError(f"CIFAR ResNet depth must be 6n+2 and >=8, got {depth}")
if base_width <= 0:
@@ -54,6 +55,13 @@ class CIFARLocalResNet:
self.n_classes = int(n_classes)
self.device = str(device)
self.dtype = dtype
+ if normalization not in ("none", "batchnorm"):
+ raise ValueError(f"unknown normalization: {normalization}")
+ self.normalization = normalization
+ self.bn_momentum = float(bn_momentum)
+ self.bn_eps = float(bn_eps)
+ if not 0.0 < self.bn_momentum <= 1.0 or self.bn_eps <= 0:
+ raise ValueError("invalid BatchNorm momentum/epsilon")
self.n_blocks = 3 * self.blocks_per_stage
self.residual_scale = (1.0 / math.sqrt(self.n_blocks)
if residual_scale is None else float(residual_scale))
@@ -64,6 +72,10 @@ class CIFARLocalResNet:
self.W = []
self.layer_specs = []
self.blocks = []
+ self.gamma = []
+ self.beta = []
+ self.running_mean = []
+ self.running_var = []
def add_conv(name, in_channels, out_channels, stride, hidden_shape,
branch_scale=1.0):
@@ -72,6 +84,13 @@ class CIFARLocalResNet:
out_channels, in_channels, 3, 3, generator=generator)
* (weight_scale * math.sqrt(2.0 / fan_in)))
self.W.append(weight.to(device=device, dtype=dtype))
+ if normalization == "batchnorm":
+ self.gamma.append(torch.ones(out_channels, device=device, dtype=dtype))
+ self.beta.append(torch.zeros(out_channels, device=device, dtype=dtype))
+ self.running_mean.append(torch.zeros(
+ out_channels, device=device, dtype=dtype))
+ self.running_var.append(torch.ones(
+ out_channels, device=device, dtype=dtype))
self.layer_specs.append(ConvLayerSpec(
name=name, stride=stride, padding=1,
hidden_shape=tuple(hidden_shape), branch_scale=float(branch_scale)))
@@ -114,6 +133,8 @@ class CIFARLocalResNet:
self.mW = [torch.zeros_like(weight) for weight in self.W]
self.mW_out = torch.zeros_like(self.W_out)
self.mb_out = torch.zeros_like(self.b_out)
+ self.mgamma = [torch.zeros_like(value) for value in self.gamma]
+ self.mbeta = [torch.zeros_like(value) for value in self.beta]
@property
def hidden_shapes(self):
@@ -126,6 +147,8 @@ class CIFARLocalResNet:
@property
def n_forward_parameters(self):
return (sum(weight.numel() for weight in self.W)
+ + sum(value.numel() for value in self.gamma)
+ + sum(value.numel() for value in self.beta)
+ self.W_out.numel() + self.b_out.numel())
@property
@@ -170,7 +193,36 @@ class CIFARLocalResNet:
f"does not match hidden value {tuple(value.shape)}")
return value + perturbation
- def forward(self, x, perturbations=None, return_cache=False):
+ def _normalize(self, index, value, training, update_stats):
+ if self.normalization == "none":
+ return value, None
+ axes = (0, 2, 3)
+ if training:
+ mean = value.mean(dim=axes)
+ variance = value.var(dim=axes, unbiased=False)
+ if update_stats:
+ with torch.no_grad():
+ count = value.numel() // value.shape[1]
+ unbiased = variance * count / max(1, count - 1)
+ self.running_mean[index].lerp_(mean.detach(), self.bn_momentum)
+ self.running_var[index].lerp_(unbiased.detach(), self.bn_momentum)
+ else:
+ mean = self.running_mean[index]
+ variance = self.running_var[index]
+ inverse_std = torch.rsqrt(variance + self.bn_eps)
+ normalized = ((value - mean[None, :, None, None])
+ * inverse_std[None, :, None, None])
+ output = (self.gamma[index][None, :, None, None] * normalized
+ + self.beta[index][None, :, None, None])
+ cache = {
+ "normalized": normalized,
+ "inverse_std": inverse_std,
+ "training": bool(training),
+ }
+ return output, cache
+
+ def forward(self, x, perturbations=None, return_cache=False, training=False,
+ update_stats=False):
if x.ndim != 4 or tuple(x.shape[1:]) != (3, 32, 32):
raise ValueError(f"expected CIFAR NCHW input, got {tuple(x.shape)}")
if perturbations is not None and len(perturbations) != self.n_hidden:
@@ -181,10 +233,12 @@ class CIFARLocalResNet:
pre = x
u = F.conv2d(pre, self.W[0], stride=1, padding=1)
- h_clean = F.relu(u)
+ normalized, norm_cache = self._normalize(0, u, training, update_stats)
+ h_clean = F.relu(normalized)
hiddens.append(h_clean)
if return_cache:
- caches.append({"pre": pre, "gate": u > 0})
+ caches.append({"pre": pre, "gate": normalized > 0,
+ "normalization": norm_cache})
h = self._inject(h_clean, perturbations, 0)
for block in self.blocks:
@@ -196,18 +250,24 @@ class CIFARLocalResNet:
pre_first = h
u_first = F.conv2d(
pre_first, self.W[first], stride=block["stride"], padding=1)
- first_clean = F.relu(u_first)
+ normalized_first, first_norm_cache = self._normalize(
+ first, u_first, training, update_stats)
+ first_clean = F.relu(normalized_first)
hiddens.append(first_clean)
if return_cache:
- caches.append({"pre": pre_first, "gate": u_first > 0})
+ caches.append({"pre": pre_first, "gate": normalized_first > 0,
+ "normalization": first_norm_cache})
first_value = self._inject(first_clean, perturbations, first)
u_second = F.conv2d(first_value, self.W[second], stride=1, padding=1)
- block_pre = shortcut + self.residual_scale * u_second
+ normalized_second, second_norm_cache = self._normalize(
+ second, u_second, training, update_stats)
+ block_pre = shortcut + self.residual_scale * normalized_second
block_clean = F.relu(block_pre)
hiddens.append(block_clean)
if return_cache:
- caches.append({"pre": first_value, "gate": block_pre > 0})
+ caches.append({"pre": first_value, "gate": block_pre > 0,
+ "normalization": second_norm_cache})
h = self._inject(block_clean, perturbations, second)
features = h.mean(dim=(2, 3))
@@ -222,6 +282,26 @@ class CIFARLocalResNet:
def logits(self, x):
return self.forward(x)["logits"]
+ def _normalization_backward(self, index, delta, cache):
+ """Local BatchNorm Jacobian-vector product and affine directions."""
+ if self.normalization == "none":
+ return delta, None, None
+ normalized = cache["normalized"]
+ gamma_direction = (delta * normalized).sum(dim=(0, 2, 3))
+ beta_direction = delta.sum(dim=(0, 2, 3))
+ scaled = delta * self.gamma[index][None, :, None, None]
+ inverse_std = cache["inverse_std"][None, :, None, None]
+ if cache["training"]:
+ count = delta.shape[0] * delta.shape[2] * delta.shape[3]
+ summed = scaled.sum(dim=(0, 2, 3), keepdim=True)
+ projected = (scaled * normalized).sum(
+ dim=(0, 2, 3), keepdim=True)
+ input_delta = (inverse_std / count) * (
+ count * scaled - summed - normalized * projected)
+ else:
+ input_delta = inverse_std * scaled
+ return input_delta, gamma_direction, beta_direction
+
def local_ascent_directions(self, teaching, output_error, forward):
"""Return forward-parameter descent directions from local signals.
@@ -238,6 +318,8 @@ class CIFARLocalResNet:
raise ValueError("local directions require a cached forward pass")
batch = output_error.shape[0]
directions = []
+ gamma_directions = []
+ beta_directions = []
with torch.no_grad():
for index, (signal, cache, spec, weight) in enumerate(zip(
teaching, caches, self.layer_specs, self.W)):
@@ -245,22 +327,34 @@ class CIFARLocalResNet:
raise ValueError(
f"teaching {index} has {tuple(signal.shape[1:])}, "
f"expected {spec.hidden_shape}")
- delta = (signal * cache["gate"].to(signal.dtype)
- * spec.branch_scale)
+ post_norm_delta = (signal * cache["gate"].to(signal.dtype)
+ * spec.branch_scale)
+ delta, gamma_direction, beta_direction = self._normalization_backward(
+ index, post_norm_delta, cache["normalization"])
direction = torch.nn.grad.conv2d_weight(
cache["pre"].detach(), weight.shape, delta.detach(),
stride=spec.stride, padding=spec.padding)
directions.append(direction / batch)
+ if gamma_direction is not None:
+ gamma_directions.append(gamma_direction / batch)
+ beta_directions.append(beta_direction / batch)
output_weight = -(output_error.t() @ forward["features"].detach()) / batch
output_bias = -output_error.mean(dim=0)
- return directions, output_weight, output_bias
+ return (directions, gamma_directions, beta_directions,
+ output_weight, output_bias)
def apply_ascent(self, directions, output_weight, output_bias, eta_hidden,
- eta_output=None, momentum=0.0, weight_decay=0.0):
+ eta_output=None, momentum=0.0, weight_decay=0.0,
+ gamma_directions=None, beta_directions=None):
"""Apply simultaneously computed directions with optional momentum."""
if len(directions) != len(self.W):
raise ValueError("one direction is required for every convolution")
eta_output = eta_hidden if eta_output is None else eta_output
+ gamma_directions = [] if gamma_directions is None else gamma_directions
+ beta_directions = [] if beta_directions is None else beta_directions
+ if self.normalization == "batchnorm" and not (
+ len(gamma_directions) == len(beta_directions) == len(self.W)):
+ raise ValueError("BatchNorm directions must cover every convolution")
with torch.no_grad():
for index, (weight, direction) in enumerate(zip(self.W, directions)):
update = direction - weight_decay * weight
@@ -268,6 +362,15 @@ class CIFARLocalResNet:
self.mW[index].mul_(momentum).add_(update)
update = self.mW[index]
weight.add_(update, alpha=eta_hidden)
+ for index, (gamma_direction, beta_direction) in enumerate(zip(
+ gamma_directions, beta_directions)):
+ if momentum:
+ self.mgamma[index].mul_(momentum).add_(gamma_direction)
+ self.mbeta[index].mul_(momentum).add_(beta_direction)
+ gamma_direction = self.mgamma[index]
+ beta_direction = self.mbeta[index]
+ self.gamma[index].add_(gamma_direction, alpha=eta_hidden)
+ self.beta[index].add_(beta_direction, alpha=eta_hidden)
out_update = output_weight - weight_decay * self.W_out
if momentum:
self.mW_out.mul_(momentum).add_(out_update)
@@ -279,18 +382,30 @@ class CIFARLocalResNet:
def bp_step(self, x, y, eta, momentum=0.0, weight_decay=0.0):
"""Exact-backprop comparator on the identical forward architecture."""
- parameters = self.W + [self.W_out, self.b_out]
+ parameters = self.W + self.gamma + self.beta + [self.W_out, self.b_out]
for parameter in parameters:
parameter.requires_grad_(True)
- loss = F.cross_entropy(self.logits(x), y)
+ loss = F.cross_entropy(
+ self.forward(x, training=True, update_stats=True)["logits"], y)
gradients = torch.autograd.grad(loss, parameters)
+ n_conv = len(self.W)
with torch.no_grad():
- conv_directions = [-gradient for gradient in gradients[:-2]]
+ conv_directions = [-gradient for gradient in gradients[:n_conv]]
+ if self.normalization == "batchnorm":
+ gamma_directions = [
+ -gradient for gradient in gradients[n_conv:2 * n_conv]]
+ beta_directions = [
+ -gradient for gradient in gradients[2 * n_conv:3 * n_conv]]
+ else:
+ gamma_directions = []
+ beta_directions = []
output_weight = -gradients[-2]
output_bias = -gradients[-1]
self.apply_ascent(
conv_directions, output_weight, output_bias, eta,
- momentum=momentum, weight_decay=weight_decay)
+ momentum=momentum, weight_decay=weight_decay,
+ gamma_directions=gamma_directions,
+ beta_directions=beta_directions)
for parameter in parameters:
parameter.requires_grad_(False)
return float(loss.detach())
@@ -425,8 +540,9 @@ def simultaneous_conv_node_perturbation(net, x, y, clean_forward, sigma=1e-2,
Independent Rademacher interventions are injected into every hidden map in
the same plus/minus evaluations. Cross-layer interference is zero mean and
- is handled by the variance theorem in ``THEORY.md``. Plus and minus trials
- are concatenated into one expanded batch for GPU efficiency.
+ is handled by the variance theorem in ``THEORY.md``. The antithetic trials
+ are evaluated as separate B-sized batches: concatenating them would couple
+ their BatchNorm statistics and change the intervention being estimated.
"""
if sigma <= 0:
raise ValueError("perturbation sigma must be positive")
@@ -439,28 +555,47 @@ def simultaneous_conv_node_perturbation(net, x, y, clean_forward, sigma=1e-2,
targets = [torch.zeros_like(hidden) for hidden in clean_forward["hiddens"]]
diagnostic_directions = []
diagnostic_derivatives = []
- expanded_x = torch.cat((x, x), dim=0)
- expanded_y = torch.cat((y, y), dim=0)
for _ in range(n_directions):
directions = []
- perturbations = []
+ plus_perturbations = []
+ minus_perturbations = []
for hidden in clean_forward["hiddens"]:
direction = torch.empty_like(hidden).bernoulli_(
0.5, generator=generator).mul_(2).sub_(1)
directions.append(direction)
- perturbations.append(torch.cat(
- (sigma * direction, -sigma * direction), dim=0))
- perturbed = net.forward(expanded_x, perturbations=perturbations)
- losses = F.cross_entropy(perturbed["logits"], expanded_y, reduction="none")
- plus, minus = losses.chunk(2)
- directional = (plus - minus) / (2.0 * sigma)
+ plus_perturbations.append(sigma * direction)
+ minus_perturbations.append(-sigma * direction)
+ plus_forward = net.forward(
+ x, perturbations=plus_perturbations,
+ training=True, update_stats=False)
+ minus_forward = net.forward(
+ x, perturbations=minus_perturbations,
+ training=True, update_stats=False)
+ plus = F.cross_entropy(plus_forward["logits"], y, reduction="none")
+ minus = F.cross_entropy(minus_forward["logits"], y, reduction="none")
+ if net.normalization == "batchnorm":
+ # BN couples examples. The per-example loss difference is not a
+ # valid node-perturbation target because ell_i also responds to
+ # xi_j for j != i. The scalar batch objective is valid; multiplying
+ # its derivative by B recovers the derivative of the summed loss,
+ # matching the per-example signal convention of the local update.
+ batch_directional = (plus.mean() - minus.mean()) / (2.0 * sigma)
+ directional = batch_directional.mul(x.shape[0]).expand(x.shape[0])
+ else:
+ batch_directional = None
+ directional = (plus - minus) / (2.0 * sigma)
for index, direction in enumerate(directions):
expand = directional.reshape(
directional.shape[0], *([1] * (direction.ndim - 1)))
targets[index].add_(-expand * direction / n_directions)
if return_diagnostics:
diagnostic_directions.append(directions)
- diagnostic_derivatives.append(directional)
+ diagnostic_derivatives.append({
+ "scaled_directional": directional,
+ "batch_mean_directional": batch_directional,
+ "coupling": ("batch_objective" if batch_directional is not None
+ else "per_example_objective"),
+ })
if return_diagnostics:
return targets, {
"directions": diagnostic_directions,
@@ -501,7 +636,8 @@ def conv_local_step(net, x, y, config, step, generator=None):
"""One DFA/learned-feedback/direct-NP minibatch update without autograd."""
config.validate()
with torch.no_grad():
- forward = net.forward(x, return_cache=True)
+ forward = net.forward(
+ x, return_cache=True, training=True, update_stats=True)
logits = forward["logits"]
loss = F.cross_entropy(logits, y)
output_error = (torch.softmax(logits, dim=1)
@@ -519,12 +655,15 @@ def conv_local_step(net, x, y, config, step, generator=None):
weight_teaching = targets if config.direct_node_perturbation else teaching
if weight_teaching is None:
raise RuntimeError("direct perturbation target is unavailable")
- directions, output_weight, output_bias = net.local_ascent_directions(
+ (directions, gamma_directions, beta_directions,
+ output_weight, output_bias) = net.local_ascent_directions(
weight_teaching, output_error, forward)
net.apply_ascent(
directions, output_weight, output_bias,
eta_hidden=config.eta, eta_output=config.eta_output,
- momentum=config.momentum, weight_decay=config.weight_decay)
+ momentum=config.momentum, weight_decay=config.weight_decay,
+ gamma_directions=gamma_directions,
+ beta_directions=beta_directions)
calibration = None
if did_perturb and config.learn_A:
calibration = net.calibrate_apical(
@@ -554,7 +693,8 @@ def conv_apical_calibration_step(net, x, y, config, generator=None):
config.validate()
if not config.learn_A:
raise ValueError("apical-only calibration requires learn_A=True")
- forward = net.forward(x, return_cache=False)
+ forward = net.forward(
+ x, return_cache=False, training=True, update_stats=False)
logits = forward["logits"]
output_error = (torch.softmax(logits, dim=1)
- F.one_hot(y, net.n_classes).to(logits.dtype))
@@ -571,7 +711,7 @@ def conv_apical_calibration_step(net, x, y, config, generator=None):
def conv_alignment_report(net, x, y, config):
"""Measure apical alignment to exact hidden gradients; never used to learn."""
- parameters = net.W + [net.W_out, net.b_out]
+ parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out]
for parameter in parameters:
parameter.requires_grad_(True)
forward = net.forward(x)