"""Minimal endogenous shared-apical-path feasibility model. The context field in this module is part of the forward computation and is required to solve the conditional task. During learning, the same apical measurement contains that ordinary field plus reciprocal KP instruction. There is no generated nuisance or bias term. """ from dataclasses import dataclass import math import torch import torch.nn.functional as F CONDITIONS = ("oracle", "raw_shared", "innovation", "matched_raw") def conditional_selector_data(n, seed, device="cpu"): """Balanced contextual selector task with no context on the basal input.""" if n % 2: raise ValueError("conditional selector data size must be even") generator = torch.Generator(device="cpu").manual_seed(seed) x = torch.randn(n, 2, generator=generator) z = torch.arange(n, dtype=torch.long).remainder(2) permutation = torch.randperm(n, generator=generator) x, z = x[permutation], z[permutation] selected = x.gather(1, z[:, None]).squeeze(1) y = (selected > 0).long() return x.to(device), z.to(device), y.to(device) @dataclass(frozen=True) class SharedFeedbackConfig: width: int = 64 hidden_layers: int = 2 learning_rate: float = 0.03 reciprocal_learning_rate: float = 0.03 momentum: float = 0.9 weight_decay: float = 1e-4 context_scale: float = 1.0 class SharedFeedbackNet: """Context-conditioned MLP with independently stored reciprocal weights.""" def __init__(self, config=SharedFeedbackConfig(), seed=3101, device="cpu", dtype=torch.float32): if config.hidden_layers < 1: raise ValueError("shared-feedback model needs a hidden population") self.config = config self.device = torch.device(device) self.dtype = dtype sizes = [2] + [config.width] * config.hidden_layers + [2] generator = torch.Generator(device="cpu").manual_seed(seed) reciprocal_generator = torch.Generator(device="cpu").manual_seed(seed + 1) context_generator = torch.Generator(device="cpu").manual_seed(seed + 2) def normal(shape, scale, source): return (torch.randn(*shape, generator=source) * scale).to( device=self.device, dtype=dtype) self.W = [normal((sizes[i + 1], sizes[i]), 1.0 / math.sqrt(sizes[i]), generator) for i in range(len(sizes) - 1)] # Q[i] corresponds to W[i] and has the same storage orientation. Q[0] # is absent because the basal input does not need a transported field. self.Q = [None] + [normal(tuple(self.W[i].shape), 1.0 / math.sqrt(sizes[i]), reciprocal_generator) for i in range(1, len(self.W))] self.C = [normal((config.width, 2), config.context_scale / math.sqrt(2.0), context_generator) for _ in range(config.hidden_layers)] self.P = [torch.zeros(config.width, device=self.device, dtype=dtype) for _ in range(config.hidden_layers)] self.P_bias = [torch.zeros_like(value) for value in self.P] self.mW = [torch.zeros_like(value) for value in self.W] self.mQ = [None] + [torch.zeros_like(value) for value in self.Q[1:]] def clone(self): copied = SharedFeedbackNet( self.config, seed=0, device=self.device, dtype=self.dtype) for name in ("W", "Q", "C", "P", "P_bias", "mW", "mQ"): source = getattr(self, name) target = [] for value in source: target.append(None if value is None else value.clone()) setattr(copied, name, target) return copied def context_fields(self, z, enabled=True): onehot = F.one_hot(z, num_classes=2).to(self.dtype) if not enabled: return [torch.zeros((z.shape[0], self.config.width), device=self.device, dtype=self.dtype) for _ in self.C] return [onehot @ projection.t() for projection in self.C] def forward(self, x, z, context_enabled=True): fields = self.context_fields(z, enabled=context_enabled) h = [x] u = [] for layer in range(self.config.hidden_layers): value = h[-1] @ self.W[layer].t() + fields[layer] u.append(value) h.append(torch.tanh(value)) logits = h[-1] @ self.W[-1].t() h.append(logits) return {"h": h, "u": u, "context": fields, "logits": logits} def predictor(self, layer, soma): return self.P[layer] * soma + self.P_bias[layer] @torch.no_grad() def fit_neutral_predictor(self, x, z): """Per-cell affine neutral fit; labels and instruction are not inputs.""" state = self.forward(x, z) reports = [] for layer, target in enumerate(state["context"]): soma = state["h"][layer + 1] soma_centered = soma - soma.mean(0) target_centered = target - target.mean(0) variance = soma_centered.square().mean(0) covariance = (soma_centered * target_centered).mean(0) slope = covariance / variance.clamp_min(1e-8) intercept = target.mean(0) - slope * soma.mean(0) self.P[layer].copy_(slope) self.P_bias[layer].copy_(intercept) prediction = slope * soma + intercept residual = target - prediction target_ss = target_centered.square().sum(0) residual_ss = residual.square().sum(0) valid = target_ss > 1e-12 r2 = 1.0 - residual_ss[valid] / target_ss[valid] reports.append({ "mean_per_cell_r2": float(r2.mean()), "context_rms": float(target.square().mean().sqrt()), "residual_context_rms_ratio": float( residual.square().mean().sqrt() / target.square().mean().sqrt().clamp_min(1e-12)), "neutral_observations": int(x.shape[0]), "instruction_observations": 0, }) return reports def select_shared_signal(net, layer, instruction, context, soma, condition): raw = instruction + context innovation = raw - net.predictor(layer, soma) if condition == "oracle": used = instruction elif condition == "raw_shared": used = raw elif condition == "innovation": used = innovation elif condition == "matched_raw": scale = (innovation.norm(dim=1, keepdim=True) / raw.norm(dim=1, keepdim=True).clamp_min(1e-12)) used = raw * scale elif condition == "exact_subtraction": used = raw - context else: raise ValueError(f"unknown shared-feedback condition: {condition}") return used, raw, innovation @torch.no_grad() def shared_feedback_step(net, x, z, y, condition): """One manual modified-KP step with a shared apical measurement.""" state = net.forward(x, z) h, u, context = state["h"], state["u"], state["context"] probabilities = torch.softmax(state["logits"], dim=1) output_instruction = F.one_hot(y, num_classes=2).to(net.dtype) - probabilities batch = x.shape[0] hidden_deltas = [None] * net.config.hidden_layers used_fields = [None] * net.config.hidden_layers raw_fields = [None] * net.config.hidden_layers innovation_fields = [None] * net.config.hidden_layers child_delta = output_instruction for layer in reversed(range(net.config.hidden_layers)): instruction = child_delta @ net.Q[layer + 1] used, raw, innovation = select_shared_signal( net, layer, instruction, context[layer], h[layer + 1], condition) delta = used * (1.0 - torch.tanh(u[layer]).square()) hidden_deltas[layer] = delta used_fields[layer] = used raw_fields[layer] = raw innovation_fields[layer] = innovation child_delta = delta directions = [hidden_deltas[0].t() @ h[0] / batch] for layer in range(1, net.config.hidden_layers): directions.append(hidden_deltas[layer].t() @ h[layer] / batch) directions.append(output_instruction.t() @ h[-2] / batch) for layer, direction in enumerate(directions): net.mW[layer].mul_(net.config.momentum).add_( direction - net.config.weight_decay * net.W[layer]) net.W[layer].add_(net.mW[layer], alpha=net.config.learning_rate) if layer > 0: net.mQ[layer].mul_(net.config.momentum).add_( direction - net.config.weight_decay * net.Q[layer]) net.Q[layer].add_( net.mQ[layer], alpha=net.config.reciprocal_learning_rate) loss = F.cross_entropy(state["logits"], y) return float(loss), { "directions": directions, "used": used_fields, "raw": raw_fields, "innovation": innovation_fields, "context": context, } @torch.no_grad() def evaluate_shared_feedback(net, x, z, y, context_enabled=True, batch_size=512): correct = 0 total_loss = 0.0 for start in range(0, x.shape[0], batch_size): stop = min(start + batch_size, x.shape[0]) logits = net.forward( x[start:stop], z[start:stop], context_enabled=context_enabled)["logits"] total_loss += float(F.cross_entropy( logits, y[start:stop], reduction="sum")) correct += int((logits.argmax(1) == y[start:stop]).sum()) return { "accuracy": correct / x.shape[0], "loss": total_loss / x.shape[0], }