summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:01:42 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:01:42 -0500
commitc0a12f4ad897f6f9f9e432d622cabdf46c40806f (patch)
treec3e8539bd46536167a556f51f05738aa7f82e232 /sdil
parent7fd707f6068da18ea1b88100892613937b060936 (diff)
oral-a: add exact local convolution eligibilities
Diffstat (limited to 'sdil')
-rw-r--r--sdil/conv.py298
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