diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:01:42 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 06:01:42 -0500 |
| commit | c0a12f4ad897f6f9f9e432d622cabdf46c40806f (patch) | |
| tree | c3e8539bd46536167a556f51f05738aa7f82e232 /sdil | |
| parent | 7fd707f6068da18ea1b88100892613937b060936 (diff) | |
oral-a: add exact local convolution eligibilities
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/conv.py | 298 |
1 files changed, 298 insertions, 0 deletions
diff --git a/sdil/conv.py b/sdil/conv.py new file mode 100644 index 0000000..5971dea --- /dev/null +++ b/sdil/conv.py @@ -0,0 +1,298 @@ +"""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 |
