summaryrefslogtreecommitdiff
path: root/src/zbp_scaling/zbp/rules.py
diff options
context:
space:
mode:
authoryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
committeryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
commit6a544fabfc2af22e4d5823410dd2387b5af89ea9 (patch)
tree0abd67bdda420deed27428b621fb59db8be07f41 /src/zbp_scaling/zbp/rules.py
scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe
Diffstat (limited to 'src/zbp_scaling/zbp/rules.py')
-rw-r--r--src/zbp_scaling/zbp/rules.py264
1 files changed, 264 insertions, 0 deletions
diff --git a/src/zbp_scaling/zbp/rules.py b/src/zbp_scaling/zbp/rules.py
new file mode 100644
index 0000000..d556201
--- /dev/null
+++ b/src/zbp_scaling/zbp/rules.py
@@ -0,0 +1,264 @@
+"""Textbook BP-free learning rules on the physical/digital partition.
+
+Every rule is expressed in the same terms as ZBP: it decides which error signal reaches each physical block's
+OUTPUT and what (if anything) is propagated to its INPUT; block parameters always update locally from the
+incoming error (in-block autograd = the physical-equivalent local rule) and digital layers use autograd.
+
+ drtp Direct Random Target Projection (Frenkel, Lefebvre, Bol 2021): DFA whose feedback signal is the
+ (negated) one-hot target instead of the output error. Handled inside ZBPFunction (mode 'drtp').
+ pepita PEPITA (Dellaferrera & Kreiman 2022): a second forward pass with the network input modulated by a
+ fixed random projection of the output error; every parameterised unit's error is its own
+ (clean - modulated) output, gradients are taken on the modulated pass, nothing propagates.
+ wm weight mirror (Akrout et al. 2019): textbook FA whose feedback matrices are learned in a mirror
+ phase from input noise / output readout of each linear layer (B -> W, so FA -> BP as B converges).
+ dtp difference target propagation (Lee et al. 2015), MLP / residual MLP: learned inverses g_l propagate
+ targets, each block minimises |f_l(h_{l-1}) - t_l|^2 locally.
+ ff forward-forward (Hinton 2022), MLP / residual MLP: per-block goodness on positive / negative data
+ (label embedded in the input), inputs length-normalised between blocks, no propagation.
+"""
+import math
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from .autograd import ZBPBlock, zbp_blocks, DFA
+from .fa import FALinear, FAConv2d
+
+
+# ----------------------------------------------------------------------------------------------- units
+def _inside(module, containers):
+ return any(m is module for m in containers)
+
+
+def learning_units(model):
+ """PEPITA units: every ZBPBlock (a physical block is one unit: only its output is observable) plus every
+ parameterised leaf module that is not inside a ZBPBlock."""
+ blocks = zbp_blocks(model)
+ inside = set()
+ for b in blocks:
+ for m in b.modules():
+ if m is not b:
+ inside.add(id(m))
+ units = []
+ for m in model.modules():
+ if isinstance(m, ZBPBlock):
+ units.append(m)
+ elif id(m) not in inside and not isinstance(m, ZBPBlock) and any(True for _ in m.parameters(recurse=False)):
+ units.append(m)
+ return units
+
+
+# ----------------------------------------------------------------------------------------------- PEPITA
+class Pepita:
+ """Holds the fixed random projection F of the output error onto the network input (or the embedding)."""
+
+ def __init__(self, n_out, in_shape, device, scale=0.05, seed=0, embed=None):
+ g = torch.Generator(device="cpu"); g.manual_seed(seed)
+ n_in = math.prod(in_shape)
+ # Dellaferrera & Kreiman: F ~ U(-a, a), a = sqrt(6 / (n_in + n_out)) * scale-ish; we use their 0.05 default
+ a = math.sqrt(6.0 / (n_in + n_out))
+ self.F = ((torch.rand(n_out, n_in, generator=g) * 2 - 1) * a * scale).to(device)
+ self.in_shape = tuple(in_shape)
+ self.embed = embed # nn.Embedding to modulate instead of the raw input (language models)
+ self._mod = None
+ if embed is not None:
+ embed.register_forward_hook(self._embed_hook)
+
+ def _embed_hook(self, m, inp, out):
+ return out if self._mod is None else out - self._mod
+
+ def step(self, model, x, y, n_out, blocks_mode="zero"):
+ """One PEPITA step: returns the clean loss. Leaves .grad on every parameter (gradient-like quantities)."""
+ units = learning_units(model)
+ blocks = zbp_blocks(model)
+ modes = [b.cfg for b in blocks]
+ for b in blocks:
+ b.cfg = b.cfg.replace(mode=blocks_mode) # no propagation through physical blocks
+ # clean pass
+ outs_clean = {}
+ hs = [u.register_forward_hook(lambda m, i, o, u=u: outs_clean.__setitem__(id(u), o.detach())) for u in units]
+ with torch.no_grad():
+ logits = model(x)
+ for h in hs:
+ h.remove()
+ if logits.dim() == 3: # LM: per-token error
+ e = logits.softmax(-1) - F.one_hot(y, n_out).to(logits.dtype)
+ e = e / (logits.shape[0] * logits.shape[1])
+ else:
+ e = (logits.softmax(-1) - F.one_hot(y, n_out).to(logits.dtype)) / logits.shape[0]
+ loss = F.cross_entropy(logits.reshape(-1, n_out), y.reshape(-1))
+ # modulated pass: input (or embedding) minus F e, each unit's input detached so nothing chains
+ if self.embed is None:
+ x_mod = x - (e.reshape(-1, n_out) @ self.F).reshape(x.shape) * x.shape[0] # undo the 1/B in e
+ else:
+ self._mod = (e @ self.F).reshape(*e.shape[:-1], -1) * (e.shape[0] * e.shape[1])
+ x_mod = x
+ outs_mod = {}
+ def pre(m, inp):
+ return tuple(t.detach().requires_grad_(t.is_floating_point()) if torch.is_tensor(t) else t for t in inp)
+ hs = [u.register_forward_pre_hook(pre) for u in units]
+ hs += [u.register_forward_hook(lambda m, i, o, u=u: outs_mod.__setitem__(id(u), o)) for u in units]
+ logits_mod = model(x_mod)
+ for h in hs:
+ h.remove()
+ self._mod = None
+ # inject: last unit (logits) gets the true error, every other unit its activation difference
+ tensors, grads = [], []
+ last = next((u for u in units if outs_mod.get(id(u)) is logits_mod), units[-1])
+ n_units = logits.shape[0] * (logits.shape[1] if logits.dim() == 3 else 1) # batch (x tokens) average, as for e
+ for u in units:
+ o = outs_mod.get(id(u))
+ if o is None or not torch.is_tensor(o) or not o.requires_grad:
+ continue
+ g = e if u is last else (outs_clean[id(u)] - o.detach()) / n_units
+ tensors.append(o); grads.append(g)
+ torch.autograd.backward(tensors, grads)
+ for b, c in zip(blocks, modes):
+ b.cfg = c
+ return loss.detach()
+
+
+# ----------------------------------------------------------------------------------------------- weight mirror
+@torch.no_grad()
+def mirror_update(model, rate=0.01, n_noise=256, gen=None):
+ """Akrout et al. 2019 mirror phase: for every FA layer, drive the linear part with zero-mean noise xi, read
+ y = W xi, and update B <- B + rate * (y^T xi / n) - rate * B, whose fixed point is B = W."""
+ for m in model.modules():
+ if isinstance(m, FALinear):
+ xi = torch.randn(n_noise, m.in_features, generator=gen, device=m.weight.device, dtype=m.weight.dtype)
+ y = xi @ m.weight.t()
+ m.B.mul_(1 - rate).add_(rate * (y.t() @ xi) / n_noise)
+ elif isinstance(m, FAConv2d):
+ k = m.in_channels * m.kernel_size[0] * m.kernel_size[1]
+ xi = torch.randn(n_noise, k, generator=gen, device=m.weight.device, dtype=m.weight.dtype)
+ y = xi @ m.weight.flatten(1).t()
+ m.B.mul_(1 - rate).add_(rate * (y.t() @ xi).reshape(m.B.shape) / n_noise)
+
+
+# ----------------------------------------------------------------------------------------------- DTP
+class DTP:
+ """Difference target propagation on a chain of ZBP blocks (MLP / residual MLP) with a digital readout.
+ Inverses g_l (one per block, residual form h + V2 tanh(V1 h)) are trained on noisy forward passes of the block."""
+
+ def __init__(self, model, device, sigma=0.1, eta_hat=0.5, lr_inv=1e-3):
+ self.blocks = [b for b in zbp_blocks(model) if b.estimate_input_grad] # blocks that need a target below them
+ self.first = [b for b in zbp_blocks(model) if not b.estimate_input_grad]
+ self.model, self.sigma, self.eta_hat = model, sigma, eta_hat
+ self.inv = nn.ModuleDict()
+ self.dims = {}
+ self.opt_inv = None
+ self.device, self.lr_inv = device, lr_inv
+
+ def _inverse(self, b, d_out, d_in):
+ k = b.name
+ if k not in self.inv:
+ g = nn.Sequential(nn.Linear(d_out, d_in), nn.Tanh(), nn.Linear(d_in, d_in)).to(self.device)
+ nn.init.zeros_(g[2].weight); nn.init.zeros_(g[2].bias)
+ self.inv[k] = g
+ self.opt_inv = torch.optim.Adam(self.inv.parameters(), lr=self.lr_inv)
+ g = self.inv[k]
+ return lambda h: h[..., :d_in] + g(h) if d_out >= d_in else g(h)
+
+ def step(self, x, y):
+ model = self.model
+ caps = {}
+ hs = [b.register_forward_hook(lambda m, i, o, b=b: caps.__setitem__(b.name, (i[0].detach(), o.detach()))) for b in zbp_blocks(model)]
+ with torch.no_grad():
+ hs_all = model(x)
+ for h in hs:
+ h.remove()
+ # digital readout: exact gradient w.r.t. its parameters and w.r.t. the top block's output
+ top = zbp_blocks(model)[-1]
+ h_top = caps[top.name][1].clone().requires_grad_(True)
+ logits = model.readout(h_top)
+ loss = F.cross_entropy(logits, y)
+ loss.backward() # readout params get .grad; h_top.grad = dL/dh_top
+ t = h_top.detach() - self.eta_hat * h_top.grad
+ # targets downward, local losses, inverse training
+ order = zbp_blocks(model)
+ inv_loss_total = 0.0
+ for b in reversed(order):
+ x_in, h_out = caps[b.name]
+ # local loss for the block: |f(x_in) - t|^2 (gradient w.r.t. block parameters only)
+ with torch.enable_grad():
+ out = b(x_in.detach())
+ l_loc = 0.5 * ((out - t) ** 2).sum() / x_in.shape[0]
+ grads = torch.autograd.grad(l_loc, [p for p in b.parameters() if p.requires_grad], allow_unused=True)
+ for p, g in zip([p for p in b.parameters() if p.requires_grad], grads):
+ if g is not None:
+ p.grad = g if p.grad is None else p.grad + g
+ if not b.estimate_input_grad:
+ break # first block: nothing below
+ g_inv = self._inverse(b, h_out.shape[-1], x_in.shape[-1])
+ # inverse training on a noisy forward pass of the block (one extra physical query)
+ with torch.no_grad():
+ xn = x_in + self.sigma * x_in.std() * torch.randn_like(x_in)
+ hn = b(xn)
+ with torch.enable_grad():
+ l_inv = ((g_inv(hn) - xn) ** 2).mean()
+ self.opt_inv.zero_grad(set_to_none=True)
+ l_inv.backward()
+ self.opt_inv.step()
+ inv_loss_total += l_inv.item()
+ with torch.no_grad(): # difference target for the block below
+ t = x_in + g_inv(t) - g_inv(h_out)
+ self.last_inv_loss = inv_loss_total
+ return loss.detach()
+
+
+# ----------------------------------------------------------------------------------------------- Forward-Forward
+class FF:
+ """Hinton 2022 on a chain of ZBP blocks: positive = input with the true label embedded in its first n_out
+ coordinates, negative = a wrong label; per block loss softplus(-(G_pos - theta)) + softplus(G_neg - theta) with
+ goodness G = mean(h^2); block inputs are length-normalised; prediction = label with the largest summed goodness
+ of all blocks after the first."""
+
+ def __init__(self, model, n_out=10, theta=2.0, label_scale=1.0):
+ self.model, self.n_out, self.theta, self.label_scale = model, n_out, theta, label_scale
+
+ def embed(self, x, labels):
+ x = x.flatten(1).clone()
+ x[:, :self.n_out] = F.one_hot(labels, self.n_out).to(x.dtype) * self.label_scale
+ return x
+
+ @staticmethod
+ def norm(h):
+ return h / (h.norm(dim=-1, keepdim=True) + 1e-4) * math.sqrt(h.shape[-1])
+
+ def goodness_chain(self, x, train=False):
+ """Returns the list of per-block goodness values; with train=True also leaves gradients (FF losses)."""
+ h = x
+ goods = []
+ for b in zbp_blocks(self.model):
+ h_in = self.norm(h).detach()
+ out = b(h_in) if train else b(h_in).detach()
+ goods.append(out.pow(2).mean(-1))
+ h = out.detach()
+ return goods
+
+ def step(self, x, y):
+ blocks = zbp_blocks(self.model)
+ x_pos = self.embed(x, y)
+ wrong = (y + torch.randint(1, self.n_out, y.shape, device=y.device)) % self.n_out
+ x_neg = self.embed(x, wrong)
+ h_pos, h_neg = x_pos, x_neg
+ total = 0.0
+ for b in blocks:
+ hp, hn = self.norm(h_pos).detach(), self.norm(h_neg).detach()
+ with torch.enable_grad():
+ op, on = b(hp), b(hn)
+ gp, gn = op.pow(2).mean(-1), on.pow(2).mean(-1)
+ loss = (F.softplus(-(gp - self.theta)) + F.softplus(gn - self.theta)).mean()
+ loss.backward()
+ total += loss.item()
+ h_pos, h_neg = op.detach(), on.detach()
+ return torch.tensor(total / len(blocks))
+
+ @torch.no_grad()
+ def predict(self, x):
+ scores = []
+ for c in range(self.n_out):
+ xc = self.embed(x, torch.full((x.shape[0],), c, device=x.device, dtype=torch.long))
+ goods = self.goodness_chain(xc)
+ scores.append(torch.stack(goods[1:] if len(goods) > 1 else goods).sum(0))
+ return torch.stack(scores, 1).argmax(1)