summaryrefslogtreecommitdiff
path: root/src/zbp_scaling/zbp/autograd.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/autograd.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/autograd.py')
-rw-r--r--src/zbp_scaling/zbp/autograd.py449
1 files changed, 449 insertions, 0 deletions
diff --git a/src/zbp_scaling/zbp/autograd.py b/src/zbp_scaling/zbp/autograd.py
new file mode 100644
index 0000000..e7ca6ec
--- /dev/null
+++ b/src/zbp_scaling/zbp/autograd.py
@@ -0,0 +1,449 @@
+"""ZBPBlock: an nn.Module wrapper whose backward pass never differentiates through
+the wrapped block for the *activation* error. The returned input gradient is
+built purely from (i) function queries of the block (or, in 'oracle' mode, a random
+projection of the exact VJP), plus (ii) the exact VJP of an optional *known* skip
+path (residual control variate). Parameter gradients are constructed locally
+from the incoming activation error v (J_theta^T v via in-block autograd).
+"""
+import contextlib
+import math
+import torch
+import torch.nn as nn
+
+from .config import ZBPConfig
+from .estimators import (oracle_vjp, zo_vjp, zo_vjp_cd_parts, make_h, control_vjp, Local,
+ causal_score_coordinate_k_vjp, causal_score_coordinate_qk_vjp)
+
+_GENERATORS = {}
+
+
+def get_generator(device, seed=None):
+ key = str(device)
+ if key not in _GENERATORS or seed is not None:
+ g = torch.Generator(device=device)
+ g.manual_seed(1234 if seed is None else seed)
+ _GENERATORS[key] = g
+ return _GENERATORS[key]
+
+
+def seed_probes(seed, device="cpu"):
+ get_generator(torch.device(device), seed)
+
+
+class DFA:
+ """Holds the output error e = dL/dlogits of the current backward pass for direct feedback alignment.
+ Training scripts call DFA.attach(logits) after the forward; blocks in 'dfa' mode read DFA.e.
+ DRTP (Frenkel et al. 2021) uses the negated one-hot target instead: set DFA.t via DFA.attach_target(y, n)."""
+ e = None
+ t = None
+
+ @staticmethod
+ def attach_target(y, n_out):
+ DFA.t = -torch.nn.functional.one_hot(y, n_out).float()
+
+ @staticmethod
+ def attach(logits):
+ if logits.requires_grad:
+ logits.register_hook(DFA._set)
+
+ @staticmethod
+ def _set(g):
+ DFA.e = g.detach()
+
+
+def _feedback_matrix(block, key, shape_in, shape_out, device, dtype):
+ """Fixed random feedback map R: shape_out -> shape_in (per feature dims), created once per block."""
+ mats = block.__dict__.setdefault("_fb", {})
+ if key not in mats:
+ g = torch.Generator(device="cpu")
+ g.manual_seed(hash((block.name, key)) % (2 ** 31))
+ d_in, d_out = math.prod(shape_in), math.prod(shape_out)
+ R = torch.randn(d_in, d_out, generator=g) / math.sqrt(d_out)
+ mats[key] = R.to(device=device, dtype=dtype)
+ return mats[key]
+
+
+class Recorder:
+ """Collects per-block backward diagnostics (exact vs estimated activation error)."""
+ _active = None
+
+ def __init__(self):
+ self.rows = []
+
+ def __enter__(self):
+ Recorder._active = self
+ return self
+
+ def __exit__(self, *a):
+ Recorder._active = None
+
+ @staticmethod
+ def active():
+ return Recorder._active
+
+ def add(self, **row):
+ self.rows.append(row)
+
+ def summary(self):
+ import collections
+ by = collections.defaultdict(list)
+ for r in self.rows:
+ by[r["name"]].append(r)
+ out = {}
+ for k, rows in by.items():
+ out[k] = {
+ "cos_mean": sum(r["cos"] for r in rows) / len(rows),
+ "relerr_mean": sum(r["relerr"] for r in rows) / len(rows),
+ "gnorm": sum(r["gnorm"] for r in rows) / len(rows),
+ "ghat_norm": sum(r["ghat_norm"] for r in rows) / len(rows),
+ "n": len(rows),
+ }
+ return out
+
+
+class ZBPFunction(torch.autograd.Function):
+ @staticmethod
+ def forward(ctx, x, block, *params):
+ # Explicit no_grad: no graph through the black box is ever built in the forward pass.
+ with torch.no_grad():
+ y = block.f(x)
+ if block.skip is not None:
+ y = y + (x if block.skip_is_identity else block.skip(x))
+ ctx.save_for_backward(x)
+ ctx.block = block
+ return y
+
+ @staticmethod
+ def backward(ctx, v):
+ (x,) = ctx.saved_tensors
+ block = ctx.block
+ cfg = block.cfg
+ f = block.f
+ nbd = block.batch_dims
+ v = v.contiguous()
+ if block.v_channel is not None: # hardware: the held error droops / picks up noise
+ v = block.v_channel(v)
+ f_params = block._f_params()
+ s_params = block._skip_params()
+ need_gx = ctx.needs_input_grad[0]
+ rec = Recorder.active()
+ want_exact = need_gx and (cfg.mode in ("exact", "zero", "noise") or rec is not None
+ or block.probe_slice is not None or block.score_probe is not None)
+ if cfg.mode == "replay":
+ want_exact = False
+ if cfg.mode in ("dfa", "opzo", "drtp") and block.skip is not None and not block.skip_is_identity and need_gx:
+ pass # DFA ignores the skip path entirely (standard DFA); skip params still get their local gradient
+
+ # --- local parameter gradients (and exact g_x if needed) via in-block autograd ---
+ g_exact = None
+ g_params = [None] * len(f_params)
+ if want_exact or (f_params and cfg.param_mode == "local"):
+ with torch.enable_grad():
+ xd = x.detach().requires_grad_(want_exact)
+ y = f(xd)
+ inputs = ([xd] if want_exact else []) + list(f_params)
+ if inputs:
+ grads = torch.autograd.grad(y, inputs, v, allow_unused=True)
+ if want_exact:
+ g_exact, grads = grads[0], grads[1:]
+ g_params = list(grads)
+
+ # --- activation error through the black-box branch ---
+ g_hat = None
+ queries = 0.0
+ if need_gx:
+ gen = get_generator(x.device)
+ if cfg.mode == "exact":
+ g_hat = g_exact
+ elif cfg.mode == "replay":
+ g_hat = None # full input error injected below
+ elif cfg.mode in ("zero", "noise"):
+ g_hat, queries = control_vjp(g_exact, cfg, gen, nbd)
+ elif cfg.mode == "fa":
+ # feedback alignment: J_f^T v replaced by a fixed random linear map of the incoming error.
+ # Sequence models ([B, T, d]): one d_in x d_out map shared across tokens (standard practice);
+ # otherwise a full map over the flattened feature dims.
+ if x.dim() == 3 and v.dim() == 3 and x.shape[1] == v.shape[1]:
+ R = _feedback_matrix(block, "fa_tok", x.shape[2:], v.shape[2:], x.device, x.dtype)
+ g_hat = (v @ R.t()) * block.fb_scale
+ else:
+ feat_in, feat_out = x.shape[nbd:], v.shape[nbd:]
+ R = _feedback_matrix(block, "fa", feat_in, feat_out, x.device, x.dtype)
+ g_hat = (v.reshape(*v.shape[:nbd], -1) @ R.t()).reshape(x.shape) * block.fb_scale
+ elif cfg.mode == "opzo":
+ # OPZO-style baseline (Xiao et al. 2024): DFA topology with a *learned* feedback matrix — the EMA of a
+ # zeroth-order estimate of the data-averaged Jacobian^T from this block's input to the logits,
+ # obtained from perturbed forward passes of the whole network (see opzo_update). Biased by design.
+ e = DFA.e
+ if e is None:
+ raise RuntimeError("DFA.e is unset: call DFA.attach(logits) after the forward pass")
+ M = block.__dict__.get("_opzo", {}).get("M")
+ if M is None:
+ g_hat = torch.zeros_like(x)
+ elif e.dim() == 3 and x.dim() >= 3 and e.shape[1] == x.shape[1]: # per-token feedback
+ g_hat = (e @ M.t()).reshape(x.shape)
+ else: # per-sample feedback
+ g_hat = (e.reshape(e.shape[0], -1) @ M.t()).reshape(x.shape)
+ elif cfg.mode in ("dfa", "drtp"):
+ # direct feedback alignment: the output error is projected straight to this block's input;
+ # DRTP: the (negated) one-hot target takes the place of the error
+ e = DFA.e if cfg.mode == "dfa" else DFA.t
+ if e is None:
+ raise RuntimeError("DFA.e / DFA.t is unset: call DFA.attach(logits) / DFA.attach_target(y, n) first")
+ e = e.to(x.dtype)
+ if e.dim() == 3 and x.dim() >= 3 and e.shape[1] == x.shape[1]: # per-token error [B, T, V]
+ R = _feedback_matrix(block, "dfa_tok", x.shape[2:], e.shape[2:], x.device, x.dtype)
+ g_hat = (e @ R.t()).reshape(x.shape) * block.fb_scale
+ else: # one error vector per sample
+ R = _feedback_matrix(block, "dfa", x.shape[1:], e.shape[1:], x.device, x.dtype)
+ g_hat = (e.reshape(e.shape[0], -1) @ R.t()).reshape(x.shape) * block.fb_scale
+ elif block.score_probe is not None:
+ # Probe the causal softmax score rows and map their VJP to K through QK^T.
+ # Other input slices retain their exact diagnostic VJPs (notably A^T delta for V).
+ spec = block.score_probe
+ g_hat = g_exact.clone()
+ a = spec["d"]
+ if spec.get("qk"): # Q and K both from the score-VJP; V exact via the known attention weights
+ gq, gk, queries = causal_score_coordinate_qk_vjp(
+ x, v, cfg, gen, spec["d"], spec["heads"], spec["dk"], spec["dv"], spec["window"])
+ g_hat[..., :a] = gq
+ else:
+ gk, queries = causal_score_coordinate_k_vjp(
+ x, v, cfg, gen, spec["d"], spec["heads"], spec["dk"], spec["dv"], spec["window"])
+ g_hat[..., a:a + spec["dk"]] = gk
+ elif block.probe_slice is not None:
+ # ZO estimate only for coordinates [a, b) of the last dim (with its own unit structure);
+ # the other coordinates keep the exact Jacobian (diagnostic: which input is hard to probe)
+ a, b = block.probe_slice
+ sbd = block.slice_batch_dims or nbd
+ xs = x[..., a:b].contiguous()
+ def f_slice(xsl):
+ k = xsl.shape[0] // x.shape[0]
+ xx = x.repeat(k, *([1] * (x.dim() - 1))) if k > 1 else x
+ return f(torch.cat([xx[..., :a], xsl, xx[..., b:]], dim=-1))
+ if cfg.mode == "oracle":
+ gs, queries = oracle_vjp(f_slice, xs, v, cfg, gen, sbd, None)
+ else:
+ h = make_h(f_slice, v, sbd, None, cfg.readout_noise, gen)
+ gs, queries = zo_vjp(h, xs, cfg, gen, sbd, None)
+ g_hat = g_exact.clone()
+ g_hat[..., a:b] = gs
+ elif cfg.mode == "oracle":
+ g_hat, queries = oracle_vjp(f, x, v, cfg, gen, nbd, block.local)
+ elif cfg.mode == "cd" and cfg.surrogate and block.local is None:
+ # learned linear control variate: g_hat = g_raw - s * (P_u g_sur - g_sur), g_sur = A v,
+ # with A and s fitted on *previous* steps' data only (unbiased conditional on the past)
+ h = make_h(f, v, nbd, None, cfg.readout_noise, gen)
+ cyc = block.stats["backward_calls"] if cfg.probe == "hadamard_cycle" else None
+ u, D, g_raw, queries = zo_vjp_cd_parts(h, x, cfg, gen, nbd, cyc)
+ g_hat = block._apply_surrogate(x, v, u, D, g_raw, nbd)
+ else:
+ fq = f if block.measure is None else (lambda xp: block.measure(f(xp)))
+ h = make_h(fq, v, nbd, block.local, cfg.readout_noise, gen)
+ cyc = block.stats["backward_calls"] if cfg.probe == "hadamard_cycle" else None
+ g_hat, queries = zo_vjp(h, x, cfg, gen, nbd, block.local, cyc)
+ block.stats["queries"] += queries
+ block.stats["backward_calls"] += 1
+ if rec is not None and g_hat is not None:
+ ge = g_exact.reshape(g_exact.shape[0], -1)
+ gh = g_hat.reshape(g_hat.shape[0], -1)
+ cos = torch.nn.functional.cosine_similarity(ge, gh, dim=1)
+ rel = (gh - ge).norm(dim=1) / ge.norm(dim=1).clamp_min(1e-30)
+ rec.add(name=block.name, cos=cos.mean().item(), relerr=rel.mean().item(),
+ gnorm=ge.norm(dim=1).mean().item(), ghat_norm=gh.norm(dim=1).mean().item(),
+ vnorm=v.reshape(v.shape[0], -1).norm(dim=1).mean().item(), queries=queries,
+ # batch-level (mean over samples) cosine, what parameter updates "see"
+ cos_batchmean=torch.nn.functional.cosine_similarity(
+ ge.mean(0, keepdim=True), gh.mean(0, keepdim=True), dim=1).item())
+
+ # --- exact skip path (residual control variate) ---
+ g_skip_params = [None] * len(s_params)
+ grad_x = None
+ if need_gx and cfg.mode == "replay":
+ grad_x = block.replay
+ if s_params:
+ with torch.enable_grad():
+ xd = x.detach()
+ s = block.skip(xd)
+ g_skip_params = list(torch.autograd.grad(s, list(s_params), v, allow_unused=True))
+ elif need_gx and cfg.mode in ("dfa", "opzo", "drtp"):
+ grad_x = g_hat # DFA/DRTP/OPZO: direct projection replaces the whole backward
+ if s_params: # skip-path parameters still get their local gradient
+ with torch.enable_grad():
+ xd = x.detach()
+ s = block.skip(xd)
+ g_skip_params = list(torch.autograd.grad(s, list(s_params), v, allow_unused=True))
+ elif need_gx:
+ if block.skip is None:
+ grad_x = g_hat
+ elif block.skip_is_identity:
+ grad_x = v + g_hat
+ else:
+ with torch.enable_grad():
+ xd = x.detach().requires_grad_(True)
+ s = block.skip(xd)
+ grads = torch.autograd.grad(s, [xd] + list(s_params), v, allow_unused=True)
+ grad_x = grads[0] + g_hat
+ g_skip_params = list(grads[1:])
+ elif s_params:
+ with torch.enable_grad():
+ xd = x.detach()
+ s = block.skip(xd)
+ g_skip_params = list(torch.autograd.grad(s, list(s_params), v, allow_unused=True))
+ if block.capture and grad_x is not None:
+ block.captured = grad_x.detach()
+ return (grad_x, None, *g_params, *g_skip_params)
+
+
+class ZBPBlock(nn.Module):
+ """y = f(x) or y = skip(x) + f(x).
+
+ The activation error through f is estimated with zeroth-order queries (or the oracle
+ projection); the skip path (identity or a cheap differentiable module) is propagated
+ exactly. With cfg.mode == 'bp' the block is ordinary autograd.
+ """
+
+ def __init__(self, f, cfg=None, skip=None, batch_dims=1, name="", estimate_input_grad=True, local=None):
+ super().__init__()
+ self.f = f
+ # local=(radius, stride): conv-like block whose output position p depends only on input
+ # positions within `radius` of stride*p -> per-position measurements (lower variance)
+ self.local = Local(*local) if isinstance(local, (tuple, list)) else local
+ self.cfg = cfg if cfg is not None else ZBPConfig()
+ self.batch_dims = batch_dims
+ self.name = name
+ self.estimate_input_grad = estimate_input_grad
+ if skip is None:
+ self.skip = None
+ self.skip_is_identity = False
+ elif skip == "identity" or isinstance(skip, nn.Identity):
+ self.skip = nn.Identity()
+ self.skip_is_identity = True
+ else:
+ self.skip = skip
+ self.skip_is_identity = False
+ self.stats = {"queries": 0.0, "backward_calls": 0}
+ self.fb_scale = 1.0 # scale of the fixed random feedback map in 'fa' / 'dfa' modes
+ self.probe_slice = None # (a, b): ZO-probe only last-dim coordinates a:b, exact elsewhere
+ self.slice_batch_dims = None
+ self.score_probe = None # metadata for causal score-coordinate K probing
+ self.capture = False # store the full input error of the next backward in .captured
+ self.captured = None
+ self.replay = None # 'replay' mode returns this tensor as the input error
+ self.measure = None # hardware measurement channel applied to y inside the query readout
+ self.v_channel = None # hardware error-transport channel (sample/hold) applied to the incoming v
+
+ def _f_params(self):
+ return [p for p in self.f.parameters() if p.requires_grad]
+
+ # ---- learned linear control variate (prior-free): A_t = EMA[g_raw v^T] EMA[v v^T]^-1, s_t from EMA[D P]/EMA[P^2] ----
+ def _apply_surrogate(self, x, v, u, D, g_raw, nbd):
+ cfg = self.cfg
+ n = u.shape[0]
+ bs = x.shape[:nbd]
+ B = math.prod(bs) if len(bs) else 1
+ xf, vf, gf = x.reshape(B, -1), v.reshape(B, -1), g_raw.reshape(B, -1)
+ uf, Df = u.reshape(n, B, -1), D.reshape(n, B)
+ din, dout = xf.shape[1], vf.shape[1]
+ st = self.__dict__.setdefault("_sur", None)
+ if st is None:
+ st = {"S_gv": torch.zeros(din, dout, device=x.device, dtype=x.dtype),
+ "S_vv": torch.zeros(dout, dout, device=x.device, dtype=x.dtype),
+ "A": None, "s": 0.0, "ema_dp": 0.0, "ema_pp": 0.0, "cnt": 0}
+ self.__dict__["_sur"] = st
+ ema = cfg.surrogate_ema
+ g_hat = gf
+ if st["A"] is not None:
+ with torch.no_grad():
+ g_sur = vf @ st["A"].t() # [B, din] surrogate A v
+ P = (uf * g_sur[None]).sum(-1) # [n, B] u_i^T g_sur (digital)
+ g_sur_raw = (uf * P[..., None]).mean(0) # same-probe projection of the surrogate
+ if st["s"] > 0:
+ g_hat = gf - st["s"] * (g_sur_raw - g_sur) # E[g_hat] = g for any A, s fixed before this step
+ dp, pp = (Df * P).sum().item(), (P * P).sum().item()
+ st["ema_dp"] = ema * st["ema_dp"] + (1 - ema) * dp
+ st["ema_pp"] = ema * st["ema_pp"] + (1 - ema) * pp
+ st["s"] = float(min(1.0, max(0.0, st["ema_dp"] / max(st["ema_pp"], 1e-30))))
+ with torch.no_grad(): # regression statistics for the NEXT steps
+ st["S_gv"] = ema * st["S_gv"] + (1 - ema) * (gf.t() @ vf) / B
+ st["S_vv"] = ema * st["S_vv"] + (1 - ema) * (vf.t() @ vf) / B
+ st["cnt"] += 1
+ if st["cnt"] >= 2:
+ ridge = cfg.surrogate_ridge * st["S_vv"].diagonal().mean().clamp_min(1e-30)
+ M = st["S_vv"] + ridge * torch.eye(dout, device=x.device, dtype=x.dtype)
+ st["A"] = torch.linalg.solve(M, st["S_gv"].t()).t() # [din, dout]: least squares of g_raw on v
+ self.stats["surrogate_s"] = st["s"]
+ return g_hat.reshape(g_raw.shape)
+
+ def _skip_params(self):
+ if self.skip is None or self.skip_is_identity:
+ return []
+ return [p for p in self.skip.parameters() if p.requires_grad]
+
+ def forward(self, x):
+ pert = self.__dict__.get("opzo_perturb")
+ if pert is not None: # OPZO estimation pass: perturb this block's input
+ alpha, gen = pert
+ z = torch.randn(x.shape, generator=gen, device=x.device, dtype=x.dtype)
+ self.__dict__.setdefault("_opzo", {"M": None})["z"] = z
+ x = x + alpha * z
+ if self.cfg.mode == "bp" or not torch.is_grad_enabled():
+ y = self.f(x)
+ if self.skip is not None:
+ y = y + (x if self.skip_is_identity else self.skip(x))
+ return y
+ if not self.estimate_input_grad:
+ x = x.detach()
+ return ZBPFunction.apply(x, self, *self._f_params(), *self._skip_params())
+
+ def extra_repr(self):
+ loc = f", local=({self.local.radius},{self.local.stride})" if self.local is not None else ""
+ return f"name={self.name}, mode={self.cfg.mode}, n={self.cfg.n_probes}, skip={'id' if self.skip_is_identity else (self.skip is not None)}{loc}"
+
+
+def zbp_blocks(model):
+ return [m for m in model.modules() if isinstance(m, ZBPBlock)]
+
+
+def opzo_update(model, forward_fn, out_clean, alpha, lam, gen):
+ """One OPZO estimation step: perturb the inputs of all 'opzo' blocks at once (alpha * Gaussian), run one extra
+ forward of the whole network, and update each block's feedback matrix
+ M <- lam * M + (1 - lam) * mean_units z (delta_logits / alpha)^T
+ (cross-block responses are zero-mean noise, as in OPZO's single noisy pass). Per-token when the logits are
+ [B, T, V] and the block input has a matching token dim, per-sample otherwise."""
+ blocks = [b for b in zbp_blocks(model) if b.cfg.mode == "opzo"]
+ if not blocks:
+ return
+ for b in blocks:
+ b.__dict__["opzo_perturb"] = (alpha, gen)
+ with torch.no_grad():
+ out_p = forward_fn()
+ for b in blocks:
+ b.__dict__["opzo_perturb"] = None
+ delta = (out_p - out_clean.detach()) / alpha
+ with torch.no_grad():
+ for b in blocks:
+ st = b.__dict__["_opzo"]
+ z = st.pop("z")
+ if delta.dim() == 3 and z.dim() >= 3 and z.shape[1] == delta.shape[1]:
+ zf = z.reshape(z.shape[0], z.shape[1], -1)
+ M_new = torch.einsum("btd,btv->dv", zf, delta) / (z.shape[0] * z.shape[1])
+ else:
+ M_new = z.reshape(z.shape[0], -1).t() @ delta.reshape(delta.shape[0], -1) / z.shape[0]
+ st["M"] = M_new if st["M"] is None else lam * st["M"] + (1 - lam) * M_new
+ b.stats["queries"] += 1.0
+
+
+def set_mode(model, cfg=None, **kw):
+ """Replace / update the estimator config of every ZBPBlock in a model."""
+ for b in zbp_blocks(model):
+ b.cfg = (cfg if cfg is not None else b.cfg).replace(**kw)
+
+
+def total_queries(model):
+ return sum(b.stats["queries"] for b in zbp_blocks(model))
+
+
+def reset_queries(model):
+ for b in zbp_blocks(model):
+ b.stats = {"queries": 0.0, "backward_calls": 0}