"""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 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 @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 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.direct_node_perturbation and self.pert_every != 1: raise ValueError("direct node perturbation requires a target every step") 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) did_perturb = ((config.learn_A or config.direct_node_perturbation) and step % config.pert_every == 0) targets = None if did_perturb: 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) calibration = None if did_perturb and config.learn_A: 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": math.sqrt(sum(float(value.square().sum()) for value in teaching) / sum(value.numel() for value in teaching)), "raw_apical_rms": math.sqrt(sum(float(value.square().sum()) for value in raw) / sum(value.numel() for value in raw)), "innovation_rms": math.sqrt( sum(float(value.square().sum()) for value in innovations) / sum(value.numel() for value in innovations)), } @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, _, _ = net.apical_components( output_error, forward["hiddens"], config.nuisance_scale, config.use_residual) 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