"""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. 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 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: """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, 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: 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 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)) 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 = [] 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): 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)) 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))) 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) 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): 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) + 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 def forward_macs_per_example(self): """Multiply-accumulates in convolutions plus the linear readout.""" total = 0 for weight, spec in zip(self.W, self.layer_specs): out_channels, in_channels, kh, kw = weight.shape _, height, width = spec.hidden_shape total += out_channels * height * width * in_channels * kh * kw total += self.W_out.numel() return int(total) @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 _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: 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) 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": normalized > 0, "normalization": norm_cache}) 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) 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": 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) 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, "normalization": second_norm_cache}) 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 _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. ``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 = [] gamma_directions = [] beta_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}") 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, 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, 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 if momentum: 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) 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.gamma + self.beta + [self.W_out, self.b_out] for parameter in parameters: parameter.requires_grad_(True) 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[: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, gamma_directions=gamma_directions, beta_directions=beta_directions) for parameter in parameters: parameter.requires_grad_(False) return float(loss.detach()) class CIFARHierarchicalFAResNet(CIFARLocalResNet): """Residual-DAG feedback alignment with independent convolutional weights. Feedback follows the actual child edges of the forward residual graph and uses locally available ReLU/BatchNorm Jacobians, but every convolutional feedback tensor is initialized independently and never reads its forward counterpart. This is a baseline and an infrastructure step for learned hierarchical dendritic feedback, not an SDIL novelty claim. """ def __init__(self, *args, feedback_seed=None, feedback_scale=1.0, **kwargs): model_seed = kwargs.get("seed", 0) super().__init__(*args, **kwargs) if feedback_scale <= 0: raise ValueError("feedback_scale must be positive") generator = torch.Generator(device="cpu").manual_seed( model_seed + 30011 if feedback_seed is None else feedback_seed) self.Q = [] for weight in self.W: fan_in = weight.shape[1] * weight.shape[2] * weight.shape[3] value = (torch.randn(weight.shape, generator=generator) * (feedback_scale * math.sqrt(2.0 / fan_in))) self.Q.append(value.to(device=weight.device, dtype=weight.dtype)) channels = self.W_out.shape[1] self.R_out = (torch.randn( channels, self.n_classes, generator=generator) * (feedback_scale / math.sqrt(channels))).to( device=self.W_out.device, dtype=self.W_out.dtype) @property def n_fixed_feedback_parameters(self): # Q[0] maps the stem to pixels and is not used for hidden credit. return (sum(value.numel() for value in self.Q[1:]) + self.R_out.numel()) @property def apical_macs_per_example(self): conv = 0 for weight, spec in zip(self.Q[1:], self.layer_specs[1:]): out_channels, in_channels, kh, kw = weight.shape _, height, width = spec.hidden_shape conv += out_channels * height * width * in_channels * kh * kw return int(conv + self.R_out.numel()) @staticmethod def _option_a_shortcut_transpose(value, in_channels, stride, output_shape): """Adjoint of the parameter-free option-A shortcut.""" out_channels = value.shape[1] if in_channels > out_channels: raise ValueError("option-A transpose cannot recover reduced channels") missing = out_channels - in_channels before = missing // 2 selected = value[:, before:before + in_channels] if stride == 1: if tuple(selected.shape) != tuple(output_shape): raise ValueError("shortcut transpose shape mismatch") return selected result = value.new_zeros(output_shape) result[:, :, ::2, ::2] = selected return result @torch.no_grad() def hierarchical_teaching(self, output_signal, forward, return_edge_contexts=False): """Propagate teaching fields through independent feedback convolutions. When requested, ``edge_contexts[l]`` is the local child field consumed by ``Q[l]`` and ``recipients[l]`` is the hidden population receiving that transposed-convolution contribution. These values expose the sufficient statistics for causal feedback calibration without reading any forward tensor. """ caches = forward.get("caches") hiddens = forward.get("hiddens") if caches is None or hiddens is None: raise ValueError("hierarchical feedback requires cached hidden states") if len(hiddens) != self.n_hidden: raise ValueError("hierarchical hidden population mismatch") teaching = [torch.zeros_like(value) for value in hiddens] edge_contexts = [None for _ in self.Q] recipients = [None for _ in self.Q] spatial = hiddens[-1].shape[2] * hiddens[-1].shape[3] teaching[-1].copy_( (output_signal @ self.R_out.t())[:, :, None, None] / spatial) for block in reversed(self.blocks): first = block["first"] second = block["second"] parent = first - 1 second_gate = caches[second]["gate"].to(teaching[second].dtype) second_delta = teaching[second] * second_gate branch_delta, _, _ = self._normalization_backward( second, second_delta * self.residual_scale, caches[second]["normalization"]) edge_contexts[second] = branch_delta recipients[second] = first teaching[first].add_(F.conv_transpose2d( branch_delta, self.Q[second], stride=1, padding=1)) teaching[parent].add_(self._option_a_shortcut_transpose( second_delta, block["in_channels"], block["stride"], hiddens[parent].shape)) first_gate = caches[first]["gate"].to(teaching[first].dtype) first_delta = teaching[first] * first_gate first_delta, _, _ = self._normalization_backward( first, first_delta, caches[first]["normalization"]) edge_contexts[first] = first_delta recipients[first] = parent teaching[parent].add_(F.conv_transpose2d( first_delta, self.Q[first], stride=block["stride"], padding=1, output_padding=block["stride"] - 1)) if return_edge_contexts: if any(value is None for value in edge_contexts[1:]): raise AssertionError("hierarchical edge context is incomplete") return teaching, edge_contexts, recipients return teaching @torch.no_grad() def hierarchical_parameter_subspace_calibration( net, x, y, clean_forward, output_signal, sigma=1e-2, n_directions=1, eta=1e-3, generator=None, return_diagnostics=False): """Calibrate the residual-DAG feedback maps with two causal queries. A Rademacher tensor is drawn in every Q/R parameter space. Each tensor is applied only to its locally available child field, producing one candidate intervention at the edge's parent population. All independent candidate fields are injected in the same antithetic pair. Multiplying the scalar loss derivative back into each random tensor gives an unbiased estimate of that edge's causal target moment; subtracting the current predicted moment is the exact local squared-field delta rule. No forward weight or reverse differentiation is used. """ if not isinstance(net, CIFARHierarchicalFAResNet): raise TypeError("hierarchical calibration requires a hierarchical net") if sigma <= 0 or n_directions < 1 or eta < 0: raise ValueError("invalid hierarchical calibration hyperparameters") if generator is None: generator = torch.Generator(device=x.device).manual_seed(0) hiddens = clean_forward.get("hiddens") if hiddens is None or len(hiddens) != net.n_hidden: raise ValueError("clean forward does not match hierarchical populations") teaching, contexts, recipients = net.hierarchical_teaching( output_signal, clean_forward, return_edge_contexts=True) batch = x.shape[0] target_q = [torch.zeros_like(value) if index else None for index, value in enumerate(net.Q)] target_r = torch.zeros_like(net.R_out) diagnostic_directions = [] diagnostic_derivatives = [] for _ in range(n_directions): random_q = [None] random_q.extend([ torch.empty_like(value).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) for value in net.Q[1:] ]) random_r = torch.empty_like(net.R_out).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) interventions = [torch.zeros_like(value) for value in hiddens] spatial_out = hiddens[-1].shape[2] * hiddens[-1].shape[3] interventions[-1].add_( (output_signal @ random_r.t())[:, :, None, None] / spatial_out) for index in range(1, len(net.Q)): spec = net.layer_specs[index] recipient = recipients[index] contribution = F.conv_transpose2d( contexts[index], random_q[index], stride=spec.stride, padding=spec.padding, output_padding=spec.stride - 1) if contribution.shape != interventions[recipient].shape: raise AssertionError("hierarchical intervention shape mismatch") interventions[recipient].add_(contribution) plus = F.cross_entropy(net.forward( x, perturbations=[sigma * value for value in interventions], training=True, update_stats=False)["logits"], y) minus = F.cross_entropy(net.forward( x, perturbations=[-sigma * value for value in interventions], training=True, update_stats=False)["logits"], y) # The parameter maps are batch-shared. Scale the mean-loss derivative # into a summed-loss hidden signal, then average its moment over B and # recipient spatial sites to keep one eta meaningful across stages. directional = (plus - minus) * batch / (2.0 * sigma) target_r.add_(random_r, alpha=-float(directional) / ( batch * n_directions)) for index in range(1, len(net.Q)): recipient = recipients[index] spatial = (hiddens[recipient].shape[2] * hiddens[recipient].shape[3]) target_q[index].add_( random_q[index], alpha=-float(directional) / ( batch * spatial * n_directions)) if return_diagnostics: diagnostic_directions.append({ "hidden": interventions, "Q": random_q, "R": random_r}) diagnostic_derivatives.append({ "scaled_directional": directional, "coupling": "summed_batch_objective", }) prediction_q = [None] for index in range(1, len(net.Q)): recipient = recipients[index] spec = net.layer_specs[index] spatial = (hiddens[recipient].shape[2] * hiddens[recipient].shape[3]) prediction_q.append(torch.nn.grad.conv2d_weight( teaching[recipient], net.Q[index].shape, contexts[index], stride=spec.stride, padding=spec.padding) / (batch * spatial)) prediction_r = (teaching[-1].mean(dim=(2, 3)).t() @ output_signal) / batch errors_q = [None] errors_q.extend([ target_q[index] - prediction_q[index] for index in range(1, len(net.Q)) ]) error_r = target_r - prediction_r for index in range(1, len(net.Q)): net.Q[index].add_(errors_q[index], alpha=eta) net.R_out.add_(error_r, alpha=eta) target_power = float(target_r.square().sum()) prediction_power = float(prediction_r.square().sum()) error_power = float(error_r.square().sum()) dot = float((target_r * prediction_r).sum()) parameters = target_r.numel() for index in range(1, len(net.Q)): target_power += float(target_q[index].square().sum()) prediction_power += float(prediction_q[index].square().sum()) error_power += float(errors_q[index].square().sum()) dot += float((target_q[index] * prediction_q[index]).sum()) parameters += target_q[index].numel() denominator = math.sqrt(target_power * prediction_power) calibration = { "calibration_mse": error_power / parameters, "target_power": target_power / parameters, "prediction_target_cosine": dot / denominator if denominator else 0.0, "parameter_update_rms": math.sqrt(error_power / parameters), } if return_diagnostics: return calibration, { "directions": diagnostic_directions, "directional_derivatives": diagnostic_derivatives, "teaching": teaching, "contexts": contexts, "recipients": recipients, "target_Q": target_q, "prediction_Q": prediction_q, "target_R": target_r, "prediction_R": prediction_r, } return calibration def conv_hierarchical_step(net, x, y, config): """One hierarchical-FA update using no reverse-mode graph or weight transport.""" config.validate() with torch.no_grad(): 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) - F.one_hot(y, net.n_classes).to(logits.dtype)) teaching = net.hierarchical_teaching(output_error, forward) total_units = sum(value.numel() for value in teaching) teaching_rms = math.sqrt( sum(float(value.square().sum()) for value in teaching) / total_units) (directions, gamma_directions, beta_directions, output_weight, output_bias) = net.local_ascent_directions( 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, gamma_directions=gamma_directions, beta_directions=beta_directions) return { "loss": float(loss), "did_perturb": False, "calibration": None, "predictor_mse": None, "teaching_rms": teaching_rms, "raw_apical_rms": teaching_rms, "innovation_rms": teaching_rms, } def conv_learned_hierarchical_step(net, x, y, config, step, generator=None): """One task update with optional causal calibration of hierarchical Q/R.""" config.validate() with torch.no_grad(): 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) - F.one_hot(y, net.n_classes).to(logits.dtype)) teaching = net.hierarchical_teaching(output_error, forward) total_units = sum(value.numel() for value in teaching) teaching_rms = math.sqrt( sum(float(value.square().sum()) for value in teaching) / total_units) did_perturb = config.learn_A and step % config.pert_every == 0 calibration = None if did_perturb: calibration = hierarchical_parameter_subspace_calibration( net, x, y, forward, output_error, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=generator) (directions, gamma_directions, beta_directions, output_weight, output_bias) = net.local_ascent_directions( 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, gamma_directions=gamma_directions, beta_directions=beta_directions) return { "loss": float(loss), "did_perturb": did_perturb, "calibration": calibration, "predictor_mse": None, "teaching_rms": teaching_rms, "raw_apical_rms": teaching_rms, "innovation_rms": teaching_rms, } def conv_hierarchical_alignment_report(net, x, y): """Audit hierarchical teaching against exact hidden gradients.""" 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, return_cache=True, training=True, update_stats=False) gradients = torch.autograd.grad( F.cross_entropy(forward["logits"], y), forward["hiddens"]) batch = x.shape[0] negative_gradients = [-batch * value.detach() for value in gradients] with torch.no_grad(): output_error = (torch.softmax(forward["logits"], dim=1) - F.one_hot(y, net.n_classes).to(forward["logits"].dtype)) teaching = net.hierarchical_teaching(output_error, forward) values = [float(F.cosine_similarity( left.flatten(1), right.flatten(1), dim=1).mean()) for left, right in zip(teaching, negative_gradients)] for parameter in parameters: parameter.requires_grad_(False) return { "normalization_state": "training_batch_stats_without_running_update", "teaching_negative_gradient_cosine": values, "raw_negative_gradient_cosine": values, "innovation_negative_gradient_cosine": values, } class CIFARSDILResNet(CIFARLocalResNet): """CIFAR local ResNet with per-unit apical vectorizers and predictors. ``spatial_template`` gives every feature unit a class-error vectorizer. ``channel_gated`` instead shares class coefficients across position and obtains spatially heterogeneous credit through a local ``tanh(h)`` gate. The latter respects convolutional translation sharing and sharply reduces feedback parameters. The predictor remains Harnett-faithful and diagonal: each unit fits its own affine soma--apical relation. """ def __init__(self, *args, a_scale=1.0, apical_seed=None, vectorizer_mode="spatial_template", **kwargs): model_seed = kwargs.get("seed", 0) super().__init__(*args, **kwargs) generator = torch.Generator(device="cpu").manual_seed( model_seed + 10007 if apical_seed is None else apical_seed) nuisance_generator = torch.Generator(device="cpu").manual_seed( model_seed + 20011 if apical_seed is None else apical_seed + 1) if vectorizer_mode not in ("spatial_template", "channel_gated"): raise ValueError(f"unknown convolutional vectorizer: {vectorizer_mode}") self.vectorizer_mode = vectorizer_mode self.A = [] self.A_gate = [] self.P = [] self.P_bias = [] self.Bnuis = [] for channels, height, width in self.hidden_shapes: units = channels * height * width # A global-average readout makes early per-unit gradients shrink # approximately as 1/(H*W). This scale keeps fixed-DFA controls # finite while learned A remains free to change its gain. std = a_scale / (height * width * math.sqrt(self.n_classes)) vectorizer_units = units if vectorizer_mode == "spatial_template" else channels self.A.append((torch.randn( vectorizer_units, self.n_classes, generator=generator) * std).to(device=self.device, dtype=self.dtype)) if vectorizer_mode == "channel_gated": self.A_gate.append(torch.zeros( channels, self.n_classes, device=self.device, dtype=self.dtype)) shape = (channels, height, width) self.P.append(torch.zeros(shape, device=self.device, dtype=self.dtype)) self.P_bias.append(torch.zeros(shape, device=self.device, dtype=self.dtype)) self.Bnuis.append(torch.exp( 0.25 * torch.randn(shape, generator=nuisance_generator) ).to(device=self.device, dtype=self.dtype)) @property def n_vectorizer_parameters(self): return (sum(value.numel() for value in self.A) + sum(value.numel() for value in self.A_gate)) @property def n_predictor_parameters(self): return (sum(value.numel() for value in self.P) + sum(value.numel() for value in self.P_bias)) @property def n_apical_parameters(self): return self.n_vectorizer_parameters + self.n_predictor_parameters @property def n_fixed_traffic_coefficients(self): return sum(value.numel() for value in self.Bnuis) @property def apical_macs_per_example(self): """MACs for projecting one class-error vector to all hidden units.""" if self.vectorizer_mode == "spatial_template": return sum(value.numel() for value in self.A) projection = sum(value.numel() for value in self.A + self.A_gate) gating = sum(math.prod(shape) for shape in self.hidden_shapes) return projection + gating def instruction(self, index, output_signal, hidden): shape = self.hidden_shapes[index] if self.vectorizer_mode == "spatial_template": return (output_signal @ self.A[index].t()).reshape( output_signal.shape[0], *shape) base = (output_signal @ self.A[index].t())[:, :, None, None] gate = (output_signal @ self.A_gate[index].t())[:, :, None, None] return base + torch.tanh(hidden) * gate def apical_components(self, output_signal, hiddens, nuisance_scale=0.0, use_residual=True): """Return teaching, raw apical, and innovation at every population.""" if len(hiddens) != self.n_hidden: raise ValueError("one somatic state is required per apical population") teaching = [] raw_apical = [] innovations = [] for index, hidden in enumerate(hiddens): instruction = self.instruction(index, output_signal, hidden) traffic = nuisance_scale * self.Bnuis[index] * hidden raw = instruction + traffic baseline = self.P[index] * hidden + self.P_bias[index] innovation = raw - baseline teaching.append(innovation if use_residual else raw) raw_apical.append(raw) innovations.append(innovation) return teaching, raw_apical, innovations @torch.no_grad() def predictor_step(self, hiddens, eta, nuisance_scale): """Neutral-period normalized LMS fit to soma-predictable traffic.""" squared_error = 0.0 units = 0 for index, hidden in enumerate(hiddens): target = nuisance_scale * self.Bnuis[index] * hidden residual = target - self.P[index] * hidden - self.P_bias[index] centered_h = hidden - hidden.mean(dim=0) centered_r = residual - residual.mean(dim=0) variance = centered_h.square().mean(dim=0) self.P[index].add_( (centered_r * centered_h).mean(dim=0) / (variance + 1e-6), alpha=eta) self.P_bias[index].add_(residual.mean(dim=0), alpha=eta) squared_error += float(residual.square().sum()) units += residual.numel() return squared_error / units @torch.no_grad() def calibrate_apical(self, output_signal, hiddens, predicted_teaching, targets, eta): """Local delta rule fitting innovation to causal perturbation targets.""" if not (len(hiddens) == len(predicted_teaching) == len(targets) == self.n_hidden): raise ValueError("calibration lists must cover every hidden population") batch = output_signal.shape[0] before_error = 0.0 target_power = 0.0 dot = 0.0 prediction_power = 0.0 for index, (hidden, prediction, target) in enumerate(zip( hiddens, predicted_teaching, targets)): error = target - prediction flat_error = error.flatten(1) if self.vectorizer_mode == "spatial_template": self.A[index].add_( flat_error.t() @ output_signal / batch, alpha=eta) else: spatial_error = error.mean(dim=(2, 3)) gated_error = (error * torch.tanh(hidden)).mean(dim=(2, 3)) self.A[index].add_( spatial_error.t() @ output_signal / batch, alpha=eta) self.A_gate[index].add_( gated_error.t() @ output_signal / batch, alpha=eta) before_error += float(error.square().sum()) target_power += float(target.square().sum()) prediction_power += float(prediction.square().sum()) dot += float((target * prediction).sum()) denominator = math.sqrt(target_power * prediction_power) return { "calibration_mse": before_error / sum( target.numel() for target in targets), "target_power": target_power / sum(target.numel() for target in targets), "prediction_target_cosine": dot / denominator if denominator else 0.0, } @torch.no_grad() def simultaneous_conv_node_perturbation(net, x, y, clean_forward, sigma=1e-2, n_directions=1, generator=None, return_diagnostics=False): """Forward-only antithetic targets for all convolutional populations. 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``. 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") if n_directions < 1: raise ValueError("n_directions must be positive") if len(clean_forward["hiddens"]) != net.n_hidden: raise ValueError("clean forward does not match network hidden populations") if generator is None: generator = torch.Generator(device=x.device).manual_seed(0) targets = [torch.zeros_like(hidden) for hidden in clean_forward["hiddens"]] diagnostic_directions = [] diagnostic_derivatives = [] for _ in range(n_directions): directions = [] for hidden in clean_forward["hiddens"]: direction = torch.empty_like(hidden).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) directions.append(direction) # Build one signed intervention at a time and retain only its scalar # losses. Keeping both complete hidden dictionaries would needlessly # double peak memory at ResNet-56. plus = F.cross_entropy(net.forward( x, perturbations=[sigma * direction for direction in directions], training=True, update_stats=False)["logits"], y, reduction="none") minus = F.cross_entropy(net.forward( x, perturbations=[-sigma * direction for direction in directions], training=True, update_stats=False)["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({ "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, "directional_derivatives": diagnostic_derivatives, } return targets @torch.no_grad() def channel_subspace_apical_calibration( net, x, y, clean_forward, output_signal, sigma=1e-2, n_directions=1, eta=1e-3, generator=None, return_diagnostics=False): """Calibrate channel-gated feedback in its representable causal subspace. The legacy estimator perturbs every spatial unit independently, estimates a full hidden target, and only then averages that target into the shared channel coefficients. With K=1, most of its variance lies outside the vectorizer's representable subspace. Here each intervention is instead ``(z_base + tanh(h) z_gate) / sqrt(2)`` with one Rademacher coefficient per example and channel. If ``D`` is the antithetic loss derivative and ``S`` is the number of spatial sites, ``-sqrt(2) D z/S`` is an unbiased estimate of the corresponding negative-gradient moment. Cross-example and cross-layer terms remain zero mean. Subtracting the predicted moments gives exactly the expected local delta-rule update that full unit targets would produce, without first estimating directions outside the feedback model's representable subspace. This is still forward-only causal calibration: it consumes the same two scalar loss queries per direction, never differentiates through the network, and updates only the local A/A_gate tensors. """ if getattr(net, "vectorizer_mode", None) != "channel_gated": raise ValueError("channel-subspace calibration requires channel_gated A") if sigma <= 0 or n_directions < 1 or eta < 0: raise ValueError("invalid channel-subspace calibration hyperparameters") if len(clean_forward["hiddens"]) != net.n_hidden: raise ValueError("clean forward does not match network hidden populations") if generator is None: generator = torch.Generator(device=x.device).manual_seed(0) batch = x.shape[0] target_base = [torch.zeros( batch, hidden.shape[1], device=hidden.device, dtype=hidden.dtype) for hidden in clean_forward["hiddens"]] target_gate = [torch.zeros_like(value) for value in target_base] diagnostic_directions = [] diagnostic_derivatives = [] inverse_sqrt_two = 1.0 / math.sqrt(2.0) for _ in range(n_directions): base_random = [] gate_random = [] directions = [] for hidden in clean_forward["hiddens"]: shape = (batch, hidden.shape[1]) base = torch.empty( shape, device=hidden.device, dtype=hidden.dtype).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) gate = torch.empty_like(base).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) direction = (base[:, :, None, None] + torch.tanh(hidden) * gate[:, :, None, None]) direction.mul_(inverse_sqrt_two) base_random.append(base) gate_random.append(gate) directions.append(direction) plus = F.cross_entropy(net.forward( x, perturbations=[sigma * value for value in directions], training=True, update_stats=False)["logits"], y, reduction="none") minus = F.cross_entropy(net.forward( x, perturbations=[-sigma * value for value in directions], training=True, update_stats=False)["logits"], y, reduction="none") if net.normalization == "batchnorm": batch_directional = (plus.mean() - minus.mean()) / (2.0 * sigma) directional = batch_directional.mul(batch).expand(batch) else: batch_directional = None directional = (plus - minus) / (2.0 * sigma) for index, (hidden, base, gate) in enumerate(zip( clean_forward["hiddens"], base_random, gate_random)): spatial = hidden.shape[2] * hidden.shape[3] scale = -math.sqrt(2.0) / (spatial * n_directions) target_base[index].add_(directional[:, None] * base, alpha=scale) target_gate[index].add_(directional[:, None] * gate, alpha=scale) if return_diagnostics: diagnostic_directions.append({ "hidden": directions, "base": base_random, "gate": gate_random}) diagnostic_derivatives.append({ "scaled_directional": directional, "batch_mean_directional": batch_directional, "coupling": ("batch_objective" if batch_directional is not None else "per_example_objective"), }) before_error = 0.0 target_power = 0.0 prediction_power = 0.0 dot = 0.0 update_power = 0.0 coefficients = 0 for index, (base_target, gate_target) in enumerate(zip( target_base, target_gate)): base_coefficient = output_signal @ net.A[index].t() gate_coefficient = output_signal @ net.A_gate[index].t() gate = torch.tanh(clean_forward["hiddens"][index]) gate_mean = gate.mean(dim=(2, 3)) gate_second_moment = gate.square().mean(dim=(2, 3)) # For prediction b + tanh(h) g, these are its inner products with # the two representable basis fields. The errors are therefore the # exact stochastic gradients of full-field squared prediction error. base_prediction = base_coefficient + gate_mean * gate_coefficient gate_prediction = (gate_mean * base_coefficient + gate_second_moment * gate_coefficient) base_error = base_target - base_prediction gate_error = gate_target - gate_prediction base_update = base_error.t() @ output_signal / batch gate_update = gate_error.t() @ output_signal / batch net.A[index].add_(base_update, alpha=eta) net.A_gate[index].add_(gate_update, alpha=eta) for prediction, target, error in ( (base_prediction, base_target, base_error), (gate_prediction, gate_target, gate_error)): before_error += float(error.square().sum()) target_power += float(target.square().sum()) prediction_power += float(prediction.square().sum()) dot += float((prediction * target).sum()) coefficients += target.numel() update_power += float(base_update.square().sum() + gate_update.square().sum()) denominator = math.sqrt(target_power * prediction_power) calibration = { "calibration_mse": before_error / coefficients, "target_power": target_power / coefficients, "prediction_target_cosine": dot / denominator if denominator else 0.0, "parameter_update_rms": math.sqrt( update_power / max(1, net.n_vectorizer_parameters)), } if return_diagnostics: return calibration, { "directions": diagnostic_directions, "directional_derivatives": diagnostic_derivatives, "target_base": target_base, "target_gate": target_gate, } return calibration @torch.no_grad() def vectorizer_subspace_apical_calibration( net, x, y, clean_forward, output_signal, sigma=1e-2, n_directions=1, eta=1e-3, generator=None, return_diagnostics=False): """Estimate the causal A/G delta rule directly in parameter space. Channel-subspace calibration first estimates one causal coefficient target per example/channel and then regresses those targets on ``output_signal``. This estimator instead draws Rademacher matrices with the exact shapes of A and A_gate. Their induced hidden intervention already contains the output context, so the antithetic scalar directly estimates the matrix moments ``mean(q c^T)`` and ``mean(q tanh(h) c^T)`` required by the local vectorizer delta rule. It uses the same two loss queries per direction and removes variance in coefficient directions that the shared A/G maps cannot represent. """ if getattr(net, "vectorizer_mode", None) != "channel_gated": raise ValueError("vectorizer-subspace calibration requires channel_gated A") if sigma <= 0 or n_directions < 1 or eta < 0: raise ValueError("invalid vectorizer-subspace calibration hyperparameters") if len(clean_forward["hiddens"]) != net.n_hidden: raise ValueError("clean forward does not match network hidden populations") if generator is None: generator = torch.Generator(device=x.device).manual_seed(0) batch = x.shape[0] target_base = [torch.zeros_like(value) for value in net.A] target_gate = [torch.zeros_like(value) for value in net.A_gate] diagnostic_directions = [] diagnostic_derivatives = [] inverse_sqrt_two = 1.0 / math.sqrt(2.0) for _ in range(n_directions): base_random = [] gate_random = [] directions = [] for hidden, base_map, gate_map in zip( clean_forward["hiddens"], net.A, net.A_gate): base = torch.empty_like(base_map).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) gate = torch.empty_like(gate_map).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) base_field = output_signal @ base.t() gate_field = output_signal @ gate.t() direction = (base_field[:, :, None, None] + torch.tanh(hidden) * gate_field[:, :, None, None]) direction.mul_(inverse_sqrt_two) base_random.append(base) gate_random.append(gate) directions.append(direction) plus = F.cross_entropy(net.forward( x, perturbations=[sigma * value for value in directions], training=True, update_stats=False)["logits"], y) minus = F.cross_entropy(net.forward( x, perturbations=[-sigma * value for value in directions], training=True, update_stats=False)["logits"], y) # A/G are shared across the minibatch, so their sufficient statistic is # the derivative of the summed loss even when examples are uncoupled. directional = (plus - minus) * batch / (2.0 * sigma) for index, (hidden, base, gate) in enumerate(zip( clean_forward["hiddens"], base_random, gate_random)): spatial = hidden.shape[2] * hidden.shape[3] scale = -math.sqrt(2.0) / ( batch * spatial * n_directions) target_base[index].add_(directional * base, alpha=scale) target_gate[index].add_(directional * gate, alpha=scale) if return_diagnostics: diagnostic_directions.append({ "hidden": directions, "base": base_random, "gate": gate_random}) diagnostic_derivatives.append({ "scaled_directional": directional, "coupling": "summed_batch_objective", }) before_error = 0.0 target_power = 0.0 prediction_power = 0.0 dot = 0.0 update_power = 0.0 coefficients = 0 for index, (base_target, gate_target) in enumerate(zip( target_base, target_gate)): base_coefficient = output_signal @ net.A[index].t() gate_coefficient = output_signal @ net.A_gate[index].t() gate = torch.tanh(clean_forward["hiddens"][index]) gate_mean = gate.mean(dim=(2, 3)) gate_second_moment = gate.square().mean(dim=(2, 3)) base_moment = base_coefficient + gate_mean * gate_coefficient gate_moment = (gate_mean * base_coefficient + gate_second_moment * gate_coefficient) base_prediction = base_moment.t() @ output_signal / batch gate_prediction = gate_moment.t() @ output_signal / batch base_error = base_target - base_prediction gate_error = gate_target - gate_prediction net.A[index].add_(base_error, alpha=eta) net.A_gate[index].add_(gate_error, alpha=eta) for prediction, target, error in ( (base_prediction, base_target, base_error), (gate_prediction, gate_target, gate_error)): before_error += float(error.square().sum()) target_power += float(target.square().sum()) prediction_power += float(prediction.square().sum()) dot += float((prediction * target).sum()) update_power += float(error.square().sum()) coefficients += target.numel() denominator = math.sqrt(target_power * prediction_power) calibration = { "calibration_mse": before_error / coefficients, "target_power": target_power / coefficients, "prediction_target_cosine": dot / denominator if denominator else 0.0, "parameter_update_rms": math.sqrt(update_power / coefficients), } if return_diagnostics: return calibration, { "directions": diagnostic_directions, "directional_derivatives": diagnostic_derivatives, "target_base": target_base, "target_gate": target_gate, } return calibration @dataclass class ConvSDILConfig: eta: float = 0.01 eta_output: float = None eta_A: float = 0.01 eta_P: float = 0.01 momentum: float = 0.9 weight_decay: float = 5e-4 learn_A: bool = True learn_P: bool = False use_residual: bool = True nuisance_scale: float = 0.0 pert_sigma: float = 1e-2 pert_every: int = 4 pert_directions: int = 1 apical_calibration_mode: str = "unit_targets" direct_node_perturbation: bool = False def validate(self): if self.eta <= 0 or (self.eta_output is not None and self.eta_output <= 0): raise ValueError("forward learning rates must be positive") if self.eta_A < 0 or self.eta_P < 0: raise ValueError("apical learning rates must be nonnegative") if self.pert_every < 1 or self.pert_directions < 1: raise ValueError("perturbation cadence/directions must be positive") if self.apical_calibration_mode not in ( "unit_targets", "channel_subspace", "vectorizer_subspace", "hierarchical_parameter_subspace"): raise ValueError("unknown apical calibration mode") if self.direct_node_perturbation and self.pert_every != 1: raise ValueError("direct node perturbation requires a target every step") if (self.direct_node_perturbation and self.apical_calibration_mode != "unit_targets"): raise ValueError("direct node perturbation requires unit targets") 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, training=True, update_stats=True) logits = forward["logits"] loss = F.cross_entropy(logits, y) output_error = (torch.softmax(logits, dim=1) - F.one_hot(y, net.n_classes).to(logits.dtype)) teaching, raw, innovations = net.apical_components( output_error, forward["hiddens"], config.nuisance_scale, config.use_residual) total_units = sum(value.numel() for value in teaching) teaching_rms = math.sqrt( sum(float(value.square().sum()) for value in teaching) / total_units) raw_apical_rms = math.sqrt( sum(float(value.square().sum()) for value in raw) / total_units) innovation_rms = math.sqrt( sum(float(value.square().sum()) for value in innovations) / total_units) del raw, innovations did_perturb = ((config.learn_A or config.direct_node_perturbation) and step % config.pert_every == 0) targets = None calibration = None if did_perturb: if config.apical_calibration_mode == "channel_subspace": calibration = channel_subspace_apical_calibration( net, x, y, forward, output_error, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=generator) elif config.apical_calibration_mode == "vectorizer_subspace": calibration = vectorizer_subspace_apical_calibration( net, x, y, forward, output_error, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=generator) else: targets = simultaneous_conv_node_perturbation( net, x, y, forward, sigma=config.pert_sigma, n_directions=config.pert_directions, generator=generator) weight_teaching = targets if config.direct_node_perturbation else teaching if weight_teaching is None: raise RuntimeError("direct perturbation target is unavailable") (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, gamma_directions=gamma_directions, beta_directions=beta_directions) if (did_perturb and config.learn_A and config.apical_calibration_mode == "unit_targets"): calibration = net.calibrate_apical( output_error, forward["hiddens"], teaching, targets, config.eta_A) predictor_mse = None if config.learn_P: predictor_mse = net.predictor_step( forward["hiddens"], config.eta_P, config.nuisance_scale) return { "loss": float(loss), "did_perturb": did_perturb, "calibration": calibration, "predictor_mse": predictor_mse, "teaching_rms": teaching_rms, "raw_apical_rms": raw_apical_rms, "innovation_rms": innovation_rms, } @torch.no_grad() def conv_apical_calibration_step(net, x, y, config, generator=None): """Fit A from one causal intervention event while forward weights stay fixed.""" config.validate() if not config.learn_A: raise ValueError("apical-only calibration requires learn_A=True") 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)) teaching, raw, innovations = net.apical_components( output_error, forward["hiddens"], config.nuisance_scale, config.use_residual) del raw, innovations if config.apical_calibration_mode == "channel_subspace": calibration = channel_subspace_apical_calibration( net, x, y, forward, output_error, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=generator) elif config.apical_calibration_mode == "vectorizer_subspace": calibration = vectorizer_subspace_apical_calibration( net, x, y, forward, output_error, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=generator) else: targets = simultaneous_conv_node_perturbation( net, x, y, forward, sigma=config.pert_sigma, n_directions=config.pert_directions, generator=generator) calibration = net.calibrate_apical( output_error, forward["hiddens"], teaching, targets, config.eta_A) return float(F.cross_entropy(logits, y)), calibration def conv_alignment_report(net, x, y, config): """Measure apical alignment to exact hidden gradients; never used to learn.""" 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, training=True, update_stats=False) gradients = torch.autograd.grad( F.cross_entropy(forward["logits"], y), forward["hiddens"]) batch = x.shape[0] negative_gradients = [-batch * gradient.detach() for gradient in gradients] with torch.no_grad(): output_error = (torch.softmax(forward["logits"], dim=1) - F.one_hot(y, net.n_classes).to(forward["logits"].dtype)) teaching, raw, innovations = net.apical_components( output_error, [value.detach() for value in forward["hiddens"]], config.nuisance_scale, config.use_residual) def cosine(left, right): left = left.flatten(1) right = right.flatten(1) return float(F.cosine_similarity(left, right, dim=1).mean()) report = { "normalization_state": "training_batch_stats_without_running_update", "teaching_negative_gradient_cosine": [ cosine(left, right) for left, right in zip(teaching, negative_gradients)], "raw_negative_gradient_cosine": [ cosine(left, right) for left, right in zip(raw, negative_gradients)], "innovation_negative_gradient_cosine": [ cosine(left, right) for left, right in zip(innovations, negative_gradients)], } for parameter in parameters: parameter.requires_grad_(False) return report @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