diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:14:42 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:14:42 -0500 |
| commit | 200625df021c02e83aeecd496b4d4f5f3ffe8ad5 (patch) | |
| tree | 126decd39e5e6fcedecba89ee698589d8a89b072 /sdil | |
| parent | bcd845b60dc85e4f46bf1ab343405c1e40ed0860 (diff) | |
oral-a: make BatchNorm credit local and causal
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/conv.py | 218 |
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) |
