diff options
Diffstat (limited to 'src/zbp_scaling/zbp/rules.py')
| -rw-r--r-- | src/zbp_scaling/zbp/rules.py | 264 |
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) |
