summaryrefslogtreecommitdiff
path: root/sdil/shared_feedback.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 14:19:26 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 14:19:26 -0500
commitb3457848820d0818840e7052c41f01c193d04a67 (patch)
treec0ad78c897787035cea3982f71ccf5b0855ecb9e /sdil/shared_feedback.py
parent20a568d28375a8477eb7f55c533ac2338756ba59 (diff)
experiment: implement shared-feedback feasibility gate
Diffstat (limited to 'sdil/shared_feedback.py')
-rw-r--r--sdil/shared_feedback.py236
1 files changed, 236 insertions, 0 deletions
diff --git a/sdil/shared_feedback.py b/sdil/shared_feedback.py
new file mode 100644
index 0000000..9a38e0f
--- /dev/null
+++ b/sdil/shared_feedback.py
@@ -0,0 +1,236 @@
+"""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],
+ }
+