"""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. Forward parameters are plain tensors. The local rule uses only the stored presynaptic activation, a postsynaptic ReLU gate, and a teaching vector at the same hidden population. ``torch.nn.grad.conv2d_weight`` evaluates their local correlation efficiently; it does not traverse a reverse-mode graph or access downstream weights. Autograd is confined to ``bp_step`` and diagnostic smoke tests for the exact BP comparator. """ from dataclasses import dataclass import math import torch import torch.nn.functional as F @dataclass(frozen=True) class ConvLayerSpec: """Static metadata for one locally updated convolution.""" name: str stride: int padding: int hidden_shape: tuple branch_scale: float class CIFARLocalResNet: """Normalization-free 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 output ReLU. Consequently there is exactly one teaching population per convolution, including a branch-scale factor for each second convolution. """ 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): 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: raise ValueError(f"base_width must be positive, got {base_width}") self.depth = int(depth) self.blocks_per_stage = (depth - 2) // 6 self.base_width = int(base_width) self.n_classes = int(n_classes) self.device = str(device) self.dtype = dtype 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)) if not self.residual_scale > 0: raise ValueError("residual_scale must be positive") generator = torch.Generator(device="cpu").manual_seed(seed) self.W = [] self.layer_specs = [] self.blocks = [] def add_conv(name, in_channels, out_channels, stride, hidden_shape, branch_scale=1.0): fan_in = 9 * in_channels weight = (torch.randn( 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)) self.layer_specs.append(ConvLayerSpec( name=name, stride=stride, padding=1, hidden_shape=tuple(hidden_shape), branch_scale=float(branch_scale))) return len(self.W) - 1 channels = base_width spatial = 32 stem = add_conv("stem", 3, channels, 1, (channels, spatial, spatial)) if stem != 0: raise AssertionError("stem must be convolution zero") for stage, out_channels in enumerate( (base_width, 2 * base_width, 4 * base_width)): for block in range(self.blocks_per_stage): stride = 2 if stage > 0 and block == 0 else 1 if stride == 2: spatial //= 2 first = add_conv( f"stage{stage + 1}.block{block + 1}.conv1", channels, out_channels, stride, (out_channels, spatial, spatial)) second = add_conv( f"stage{stage + 1}.block{block + 1}.conv2", out_channels, out_channels, 1, (out_channels, spatial, spatial), self.residual_scale) self.blocks.append({ "first": first, "second": second, "in_channels": channels, "out_channels": out_channels, "stride": stride, }) channels = out_channels if len(self.W) != depth - 1: raise AssertionError( f"expected {depth - 1} convolutions, constructed {len(self.W)}") self.W_out = (torch.randn(n_classes, channels, generator=generator) / math.sqrt(channels)).to(device=device, dtype=dtype) self.b_out = torch.zeros(n_classes, device=device, dtype=dtype) 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) @property def hidden_shapes(self): return [spec.hidden_shape for spec in self.layer_specs] @property def n_hidden(self): return len(self.layer_specs) @property def n_forward_parameters(self): return (sum(weight.numel() for weight in self.W) + self.W_out.numel() + self.b_out.numel()) @staticmethod def _option_a_shortcut(x, out_channels, stride): """Original CIFAR ResNet identity shortcut with striding/zero padding.""" if stride == 2: x = x[:, :, ::2, ::2] in_channels = x.shape[1] if in_channels == out_channels: return x if in_channels > out_channels: raise ValueError("option-A shortcut cannot reduce channel count") missing = out_channels - in_channels before = missing // 2 after = missing - before chunks = [] if before: chunks.append(x.new_zeros(x.shape[0], before, x.shape[2], x.shape[3])) chunks.append(x) if after: chunks.append(x.new_zeros(x.shape[0], after, x.shape[2], x.shape[3])) return torch.cat(chunks, dim=1) def _inject(self, value, perturbations, index): if perturbations is None: return value perturbation = perturbations[index] if tuple(perturbation.shape) != tuple(value.shape): raise ValueError( f"perturbation {index} shape {tuple(perturbation.shape)} " f"does not match hidden value {tuple(value.shape)}") return value + perturbation def forward(self, x, perturbations=None, return_cache=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: raise ValueError( f"expected {self.n_hidden} perturbations, got {len(perturbations)}") hiddens = [] caches = [] pre = x u = F.conv2d(pre, self.W[0], stride=1, padding=1) h_clean = F.relu(u) hiddens.append(h_clean) if return_cache: caches.append({"pre": pre, "gate": u > 0}) h = self._inject(h_clean, perturbations, 0) for block in self.blocks: first = block["first"] second = block["second"] shortcut = self._option_a_shortcut( h, block["out_channels"], block["stride"]) pre_first = h u_first = F.conv2d( pre_first, self.W[first], stride=block["stride"], padding=1) first_clean = F.relu(u_first) hiddens.append(first_clean) if return_cache: caches.append({"pre": pre_first, "gate": u_first > 0}) 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 block_clean = F.relu(block_pre) hiddens.append(block_clean) if return_cache: caches.append({"pre": first_value, "gate": block_pre > 0}) h = self._inject(block_clean, perturbations, second) features = h.mean(dim=(2, 3)) logits = features @ self.W_out.t() + self.b_out result = {"logits": logits, "features": features, "hiddens": hiddens} if return_cache: if len(caches) != self.n_hidden: raise AssertionError("cache/hidden layer mismatch") result["caches"] = caches return result def logits(self, x): return self.forward(x)["logits"] def local_ascent_directions(self, teaching, output_error, forward): """Return forward-parameter descent directions from local signals. ``teaching[l][i]`` represents the per-example ``-d ell_i/dh_l``. Each convolutional direction averages the exact local Jacobian-vector products using only that population's cache. The output error is the per-example ``d ell_i/dlogits`` and therefore receives an explicit minus sign. """ if len(teaching) != self.n_hidden: raise ValueError(f"expected {self.n_hidden} teaching tensors") caches = forward.get("caches") if caches is None: raise ValueError("local directions require a cached forward pass") batch = output_error.shape[0] directions = [] with torch.no_grad(): for index, (signal, cache, spec, weight) in enumerate(zip( teaching, caches, self.layer_specs, self.W)): if tuple(signal.shape[1:]) != spec.hidden_shape: 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) direction = torch.nn.grad.conv2d_weight( cache["pre"].detach(), weight.shape, delta.detach(), stride=spec.stride, padding=spec.padding) directions.append(direction / batch) output_weight = -(output_error.t() @ forward["features"].detach()) / batch output_bias = -output_error.mean(dim=0) return 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): """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 with torch.no_grad(): for index, (weight, direction) in enumerate(zip(self.W, directions)): update = direction - weight_decay * weight if momentum: self.mW[index].mul_(momentum).add_(update) update = self.mW[index] weight.add_(update, alpha=eta_hidden) out_update = output_weight - weight_decay * self.W_out if momentum: self.mW_out.mul_(momentum).add_(out_update) self.mb_out.mul_(momentum).add_(output_bias) out_update = self.mW_out output_bias = self.mb_out self.W_out.add_(out_update, alpha=eta_output) self.b_out.add_(output_bias, alpha=eta_output) 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] for parameter in parameters: parameter.requires_grad_(True) loss = F.cross_entropy(self.logits(x), y) gradients = torch.autograd.grad(loss, parameters) with torch.no_grad(): conv_directions = [-gradient for gradient in gradients[:-2]] output_weight = -gradients[-2] output_bias = -gradients[-1] self.apply_ascent( conv_directions, output_weight, output_bias, eta, momentum=momentum, weight_decay=weight_decay) for parameter in parameters: parameter.requires_grad_(False) return float(loss.detach()) @torch.no_grad() def evaluate_conv(net, loader): correct = 0 total = 0 total_loss = 0.0 for x, y in loader: logits = net.logits(x) total_loss += F.cross_entropy(logits, y, reduction="sum").item() correct += (logits.argmax(dim=1) == y).sum().item() total += y.numel() return correct / total, total_loss / total