diff options
| -rw-r--r-- | experiments/baseline_run.py | 241 | ||||
| -rw-r--r-- | experiments/baseline_smoke.py | 107 | ||||
| -rwxr-xr-x | experiments/baseline_sweep.sh | 37 | ||||
| -rw-r--r-- | sdil/local_baselines.py | 241 |
4 files changed, 546 insertions, 80 deletions
diff --git a/experiments/baseline_run.py b/experiments/baseline_run.py new file mode 100644 index 0000000..a96e24d --- /dev/null +++ b/experiments/baseline_run.py @@ -0,0 +1,241 @@ +"""Canonical runners for non-backprop baselines outside the SDIL family. + +The methods intentionally do not all share one architecture: Forward-Forward +has goodness layers instead of a classifier readout, while Equilibrium +Propagation is an energy-based recurrent network with symmetric connections. +Forcing either into SDILNet would make the comparison look uniform while +silently changing the algorithm. Every JSON result records the protocol. +""" +import argparse +import json +import os +import subprocess +import sys +import time + +import torch +import torch.nn.functional as F + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.data import get_dataset, onehot +from sdil.local_baselines import FANet, PEPITANet, FFNet, EPNet + + +METHOD_SOURCES = { + "fa": { + "paper": "https://www.nature.com/articles/ncomms13276", + "protocol": "fixed random sequential feedback matrices", + }, + "pepita": { + "paper": "https://proceedings.mlr.press/v162/dellaferrera22a.html", + "code": "https://github.com/GiorgiaD/PEPITA", + "protocol": "ERIN two-forward-pass rule; He-uniform; ReLU; F scale 0.05", + }, + "ff": { + "paper": "https://www.cs.toronto.edu/~hinton/absps/FFXfinal.pdf", + "reference": "https://github.com/mpezeshki/pytorch_forward_forward", + "protocol": "greedy layerwise goodness; input length normalization; Adam", + }, + "ep": { + "paper": "https://www.frontiersin.org/journals/computational-neuroscience/articles/10.3389/fncom.2017.00024/full", + "code": "https://github.com/bscellier/Towards-a-Biologically-Plausible-Backprop", + "protocol": "two-phase energy-gradient dynamics; random beta sign", + }, +} + +_MNIST_STATS = {"mnist": (0.1307, 0.3081), "fmnist": (0.2860, 0.3530)} +_CIFAR_MEAN = torch.tensor([0.4914, 0.4822, 0.4465]) +_CIFAR_STD = torch.tensor([0.2470, 0.2435, 0.2616]) + + +def code_provenance(): + root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + try: + commit = subprocess.run(["git", "rev-parse", "HEAD"], cwd=root, check=True, + capture_output=True, text=True).stdout.strip() + dirty = bool(subprocess.run( + ["git", "status", "--porcelain", "--untracked-files=no"], cwd=root, + check=True, capture_output=True, text=True).stdout.strip()) + return {"git_commit": commit, "git_dirty": dirty} + except (OSError, subprocess.CalledProcessError): + return {"git_commit": None, "git_dirty": None} + + +def canonical_input(x, dataset, method, enabled=True): + """Restore [0,1] pixels where the authors' protocol requires them.""" + if not enabled or method not in ("pepita", "ep"): + return x + if dataset in _MNIST_STATS: + mean, std = _MNIST_STATS[dataset] + return (x * std + mean).clamp(0, 1) + if dataset == "cifar10": + mean = _CIFAR_MEAN.to(x.device, x.dtype).view(1, 3, 1) + std = _CIFAR_STD.to(x.device, x.dtype).view(1, 3, 1) + return (x.view(-1, 3, 1024) * std + mean).clamp(0, 1).view_as(x) + return x + + +@torch.no_grad() +def evaluate_method(net, loader, method, dataset, canonical): + correct = total = 0 + loss_sum = 0.0 + for x, y in loader: + x = canonical_input(x, dataset, method, canonical) + if method == "ff": + pred = net.predict(x) + elif method == "ep": + state = net._settle(x, beta=0.0, T=net.T_free) + logits = state[-1] + pred = logits.argmax(1) + loss_sum += F.mse_loss(logits, onehot(y, 10), reduction="sum").item() + else: + logits = net.logits(x) + pred = logits.argmax(1) + loss_sum += F.cross_entropy(logits, y, reduction="sum").item() + correct += (pred == y).sum().item() + total += y.shape[0] + return correct / total, loss_sum / total if method != "ff" else float("nan") + + +def ep_learning_rates(depth, base_eta=None): + """Layerwise rates from the authors' one/two/three-hidden-layer runs.""" + canonical = { + 1: [0.1, 0.05], + 2: [0.4, 0.1, 0.01], + 3: [0.128, 0.032, 0.008, 0.002], + } + if base_eta is not None: + return [base_eta / (4 ** l) for l in range(depth + 1)] + if depth in canonical: + return canonical[depth] + # The original implementation already required rapidly shrinking rates as + # depth grew. Continue its depth-3 geometric schedule for an explicit, + # logged extrapolation rather than pretending there is a canonical one. + return [0.128 / (4 ** l) for l in range(depth + 1)] + + +def build(args, n_in, device): + sizes = [n_in] + [args.width] * args.depth + [10] + if args.method == "fa": + return FANet(sizes, act=args.act, device=device, seed=args.seed, + residual=bool(args.residual), b_scale=args.feedback_scale) + if args.method == "pepita": + return PEPITANet(sizes, act="relu", device=device, seed=args.seed, + f_scale=args.pepita_f_scale, keep_prob=args.keep_prob) + if args.method == "ff": + # FF has only goodness layers; a 10-unit SDIL-style output is not part + # of the original supervised algorithm. + return FFNet([n_in] + [args.width] * args.depth, device=device, + seed=args.seed, threshold=args.ff_threshold) + if args.method == "ep": + return EPNet(sizes, device=device, seed=args.seed, beta=args.ep_beta, + dt=args.ep_dt, T_free=args.ep_free_steps, + T_nudge=args.ep_nudge_steps, random_beta_sign=True) + raise ValueError(args.method) + + +def train(args): + torch.manual_seed(args.seed) + device = args.device + train_loader, test_loader, n_in, n_out = get_dataset( + args.dataset, args.batch_size, device=device) + net = build(args, n_in, device) + canonical = bool(args.canonical_preprocess) + log = { + "args": vars(args), + "method_source": METHOD_SOURCES[args.method], + "provenance": code_provenance(), + "steps": [], + "final": {}, + } + t0 = time.time() + + if args.method == "ff": + eta = 0.03 if args.eta is None else args.eta + for layer in range(args.depth): + for epoch in range(args.epochs): + batches = 0 + for x, y in train_loader: + loss = net.train_layer(layer, x, y, eta) + batches += 1 + if args.max_batches and batches >= args.max_batches: + break + print(f"[{args.tag}] layer {layer} epoch {epoch} local_loss {loss:.4f}", + flush=True) + acc, test_loss = evaluate_method(net, test_loader, args.method, + args.dataset, canonical) + log["steps"].append({"layer_end": layer, "test_acc": acc, + "local_loss": loss}) + print(f"[{args.tag}] layer {layer} test_acc {acc:.4f}", flush=True) + else: + eta = args.eta + if eta is None: + eta = 0.05 if args.method == "fa" else (0.1 if args.method == "pepita" else None) + ep_etas = ep_learning_rates(args.depth, eta) if args.method == "ep" else None + for epoch in range(args.epochs): + if args.method == "pepita" and epoch in (60, 90): + eta *= 0.1 + batches = 0 + for x, y in train_loader: + x = canonical_input(x, args.dataset, args.method, canonical) + yoh = onehot(y, n_out, device=device) + if args.method == "fa": + loss = net.fa_step(x, y, yoh, eta, args.momentum) + elif args.method == "pepita": + loss = net.pepita_step(x, y, yoh, eta, args.momentum) + else: + loss = net.train_step(x, y, yoh, ep_etas) + batches += 1 + if args.max_batches and batches >= args.max_batches: + break + acc, test_loss = evaluate_method(net, test_loader, args.method, + args.dataset, canonical) + log["steps"].append({"epoch_end": epoch, "test_acc": acc, + "test_loss": test_loss, "train_loss": loss}) + print(f"[{args.tag}] epoch {epoch} loss {loss:.4f} test_acc {acc:.4f}", + flush=True) + + acc, test_loss = evaluate_method(net, test_loader, args.method, + args.dataset, canonical) + log["final"] = {"test_acc": acc, "test_loss": test_loss, + "wall_s": time.time() - t0} + os.makedirs(args.outdir, exist_ok=True) + path = os.path.join(args.outdir, f"{args.tag}.json") + with open(path, "w") as f: + json.dump(log, f) + print(f"[{args.tag}] DONE test_acc={acc:.4f} -> {path}", flush=True) + return log + + +def get_args(): + p = argparse.ArgumentParser() + p.add_argument("--method", required=True, choices=["fa", "pepita", "ff", "ep"]) + p.add_argument("--dataset", default="mnist", choices=["mnist", "fmnist", "cifar10"]) + p.add_argument("--depth", type=int, default=2) + p.add_argument("--width", type=int, default=500) + p.add_argument("--epochs", type=int, default=10, + help="FF: epochs per greedy layer; other methods: global epochs") + p.add_argument("--batch_size", type=int, default=64) + p.add_argument("--max_batches", type=int, default=0) + p.add_argument("--eta", type=float, default=None) + p.add_argument("--momentum", type=float, default=0.9) + p.add_argument("--act", default="tanh", choices=["tanh", "gelu", "silu", "relu"]) + p.add_argument("--residual", type=int, default=0) + p.add_argument("--feedback_scale", type=float, default=1.0) + p.add_argument("--pepita_f_scale", type=float, default=0.05) + p.add_argument("--keep_prob", type=float, default=0.9) + p.add_argument("--ff_threshold", type=float, default=2.0) + p.add_argument("--ep_beta", type=float, default=0.5) + p.add_argument("--ep_dt", type=float, default=0.5) + p.add_argument("--ep_free_steps", type=int, default=20) + p.add_argument("--ep_nudge_steps", type=int, default=4) + p.add_argument("--canonical_preprocess", type=int, default=1) + p.add_argument("--seed", type=int, default=0) + p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + p.add_argument("--outdir", default="results") + p.add_argument("--tag", default="baseline") + return p.parse_args() + + +if __name__ == "__main__": + train(get_args()) diff --git a/experiments/baseline_smoke.py b/experiments/baseline_smoke.py new file mode 100644 index 0000000..4bc3d85 --- /dev/null +++ b/experiments/baseline_smoke.py @@ -0,0 +1,107 @@ +"""Mechanism checks for the publication baselines (CPU, no dataset needed).""" +import os +import sys + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.data import onehot +from sdil.local_baselines import FANet, PEPITANet, FFNet, EPNet + + +def check_fa_residual_transport(): + torch.manual_seed(1) + x = torch.randn(32, 8) + y = torch.randint(0, 3, (32,)) + yoh = onehot(y, 3) + residual = FANet([8, 6, 6, 6, 3], seed=2, residual=True) + plain = FANet([8, 6, 6, 6, 3], seed=2, residual=False) + for net in (residual, plain): + for l in (1, 2): + net.B[l].zero_() + before_res = [w.clone() for w in residual.W] + before_plain = [w.clone() for w in plain.W] + residual.fa_step(x, y, yoh, eta=0.01) + plain.fa_step(x, y, yoh, eta=0.01) + res_changes = [(a - b).norm().item() for a, b in zip(residual.W, before_res)] + plain_changes = [(a - b).norm().item() for a, b in zip(plain.W, before_plain)] + assert all(v > 0 for v in res_changes) + assert plain_changes[0] == 0 and plain_changes[1] == 0 + assert plain_changes[2] > 0 and plain_changes[3] > 0 + print("FA residual identity transport:", res_changes) + + +def check_pepita_output_rule(): + torch.manual_seed(2) + net = PEPITANet([5, 4, 3], act="relu", seed=3, keep_prob=1.0) + x = torch.rand(16, 5) + y = torch.randint(0, 3, (16,)) + yoh = onehot(y, 3) + with torch.no_grad(): + clean = net._pepita_forward(x, None) + error = torch.softmax(clean[-1], 1) - yoh + mod = net._pepita_forward(x + error @ net.Fproj.t(), None) + mod_error = torch.softmax(mod[-1], 1) - yoh + expected = -(mod_error.t() @ mod[-2]) / x.shape[0] + old = net.W[-1].clone() + net.pepita_step(x, y, yoh, eta=0.1, momentum=0.0) + assert torch.allclose(net.W[-1] - old, 0.1 * expected, atol=1e-7, rtol=1e-5) + assert net.Fproj.abs().max() <= (6.0 / 5) ** 0.5 * 0.05 + 1e-7 + print("PEPITA modulated-output delta: exact") + + +def ff_loss(net, layer, x, y, neg): + hp = net._inputs_to_layer(layer, net._overlay(x, y)) + hn = net._inputs_to_layer(layer, net._overlay(x, neg)) + gp = net._layer_forward(layer, hp).pow(2).mean(1) + gn = net._layer_forward(layer, hn).pow(2).mean(1) + return (torch.nn.functional.softplus(-gp + net.thr) + + torch.nn.functional.softplus(gn - net.thr)).mean().item() + + +def check_ff_local_optimization(): + torch.manual_seed(3) + net = FFNet([20, 32, 32], seed=4) + x = torch.randn(128, 20) + y = torch.randint(0, 10, (128,)) + neg = (y + 1) % 10 + initial = ff_loss(net, 0, x, y, neg) + for _ in range(80): + net.train_layer(0, x, y, eta=0.03, negative_labels=neg) + final = ff_loss(net, 0, x, y, neg) + assert final < initial - 0.1 + assert net._layer_forward(0, x).shape == (128, 32) + print(f"FF layer-local loss: {initial:.4f} -> {final:.4f}") + + +def check_ep_energy_dynamics(): + torch.manual_seed(4) + net = EPNet([4, 3, 2], seed=5, dt=0.2, random_beta_sign=False) + x = torch.rand(7, 4) + y = onehot(torch.randint(0, 2, (7,)), 2) + s = [torch.rand(7, 3) * 0.8 + 0.1, torch.rand(7, 2) * 0.8 + 0.1] + beta = 0.3 + manual = [] + for k in range(net.L): + below = x if k == 0 else s[k - 1] + drive = -s[k] + below @ net.W[k].t() + net.b[k] + if k < net.L - 1: + drive = drive + s[k + 1] @ net.W[k + 1] + else: + drive = drive + 2 * beta * (y - s[k]) + manual.append((s[k] + net.dt * drive).clamp(0, 1)) + actual = net._settle(x, y, beta=beta, s=[v.clone() for v in s], T=1) + assert all(torch.allclose(a, b, atol=1e-7) for a, b in zip(actual, manual)) + old = [w.clone() for w in net.W] + net.train_step(x, y.argmax(1), y, eta=[0.01, 0.005]) + assert all(torch.isfinite(w).all() for w in net.W) + assert any(not torch.equal(a, b) for a, b in zip(net.W, old)) + print("EP one-step -d(E+beta*C)/ds dynamics: exact") + + +if __name__ == "__main__": + check_fa_residual_transport() + check_pepita_output_rule() + check_ff_local_optimization() + check_ep_energy_dynamics() + print("ALL BASELINE SMOKE CHECKS PASSED") diff --git a/experiments/baseline_sweep.sh b/experiments/baseline_sweep.sh new file mode 100755 index 0000000..bd3ff7b --- /dev/null +++ b/experiments/baseline_sweep.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +# Run canonical non-backprop baselines. Method-specific knobs can be appended. +# Usage: baseline_sweep.sh <gpu> "<methods>" "<seeds>" <dataset> <depth> <width> <epochs> <prefix> [extra args...] +set -eu + +cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +GPU="${1:?GPU index required}" +METHODS="${2:-fa pepita ff ep}" +SEEDS="${3:-0}" +DATASET="${4:-mnist}" +DEPTH="${5:-2}" +WIDTH="${6:-500}" +EPOCHS="${7:-10}" +PREFIX="${8:-canonical}" +shift 8 || true + +export CUDA_VISIBLE_DEVICES="$GPU" +export OMP_NUM_THREADS=2 +mkdir -p results logs/baselines + +for seed in $SEEDS; do + for method in $METHODS; do + tag="${PREFIX}_${DATASET}_${method}_w${WIDTH}_d${DEPTH}_s${seed}" + result="results/${tag}.json" + log="logs/baselines/${tag}.log" + if [[ -f "$result" ]]; then + echo "skip $tag (result exists)" + continue + fi + echo ">>> $tag $(date --iso-8601=seconds) gpu=$GPU" + "$PY" experiments/baseline_run.py --method "$method" --dataset "$DATASET" \ + --depth "$DEPTH" --width "$WIDTH" --epochs "$EPOCHS" --seed "$seed" \ + --tag "$tag" --outdir results "$@" > "$log" 2>&1 + grep -h DONE "$log" + done +done diff --git a/sdil/local_baselines.py b/sdil/local_baselines.py index a6a11a3..23b1b5c 100644 --- a/sdil/local_baselines.py +++ b/sdil/local_baselines.py @@ -1,6 +1,7 @@ """ -Other biologically-motivated / non-backprop local learning baselines, sharing -SDILNet's architecture and initialisation for a fair comparison: +Other biologically-motivated / non-backprop local learning baselines. FA and +PEPITA reuse SDILNet's tensor layout; FF and EP retain their native model +classes because their state spaces and objectives are fundamentally different: - FANet : Feedback Alignment (Lillicrap 2016) -- backprop through FIXED random feedback matrices instead of W^T (sequential, layer-wise). DFA is the @@ -12,7 +13,10 @@ SDILNet's architecture and initialisation for a fair comparison: it on negative (wrong label) data; inference picks max total goodness. All train with no global backward graph. Autograd, where used (FF's per-layer -local loss), never crosses layer boundaries -> the update stays local. +local loss), never crosses layer boundaries -> the update stays local. The +implementations follow the authors' published algorithms and reference code; +their deliberately different architectures are exposed rather than hidden +behind an allegedly apples-to-apples SDILNet wrapper. """ import math import torch @@ -25,9 +29,11 @@ from .core import SDILNet, ACTS # Feedback Alignment (sequential, layer-wise random feedback) # -------------------------------------------------------------------------- class FANet(SDILNet): - def __init__(self, *args, b_scale=1.0, **kw): + def __init__(self, *args, b_scale=1.0, feedback_seed=None, **kw): + model_seed = kw.get("seed", 0) super().__init__(*args, **kw) - g = torch.Generator(device="cpu").manual_seed(4242) + g = torch.Generator(device="cpu").manual_seed( + model_seed + 4242 if feedback_seed is None else feedback_seed) # fixed random feedback B[l] with the same shape as W[l], for l=1..L-1 self.B = [None] for l in range(1, self.L): @@ -39,16 +45,22 @@ class FANet(SDILNet): with torch.no_grad(): fwd = self.forward(x) h, u = fwd["h"], fwd["u"] - B = x.shape[0] e = torch.softmax(h[-1], 1) - yoh # grad wrt logits loss = F.cross_entropy(h[-1], y).item() # output layer (exact local) - delta = e # grad wrt u[L-1] - self._upd(self.L - 1, delta, h[self.L - 1], eta, momentum) - # hidden layers: propagate with fixed random B + self._upd(self.L - 1, e, h[self.L - 1], eta, momentum) + + # dh is the error wrt the hidden *state*. On a residual block, + # h'=h+alpha*phi(Wh), the identity branch transports dh exactly and + # only the nonlinear branch uses a fixed random feedback matrix. + dh = e @ self.B[self.L - 1] for l in range(self.L - 2, -1, -1): - delta = (delta @ self.B[l + 1]) * self.act_prime(u[l]) # grad wrt u[l] - self._upd(l, delta, h[l], eta, momentum) + branch_scale = self.res_alpha if self.residual and l >= 1 else 1.0 + du = branch_scale * dh * self.act_prime(u[l]) + self._upd(l, du, h[l], eta, momentum) + if l > 0: + feedback = du @ self.B[l] + dh = dh + feedback if self.residual and l >= 1 else feedback return loss def _upd(self, l, delta, pre, eta, momentum): @@ -65,39 +77,85 @@ class FANet(SDILNet): # PEPITA (forward-only, error-modulated second pass) # -------------------------------------------------------------------------- class PEPITANet(SDILNet): - def __init__(self, *args, f_scale=0.1, **kw): + """Original two-forward-pass PEPITA/ERIN rule. + + Dellaferrera & Kreiman initialise both feedforward weights and the input + error projection with He-uniform samples, then multiply the latter by 0.05. + Their fully-connected implementation has ReLU hidden units, a softmax + output, no biases, and optionally shares one dropout mask across the two + presentations. We keep logits internally but form exactly the same + softmax errors. + """ + + def __init__(self, *args, f_scale=0.05, original_init=True, + use_bias=False, keep_prob=1.0, **kw): + model_seed = kw.get("seed", 0) super().__init__(*args, **kw) - g = torch.Generator(device="cpu").manual_seed(7777) + g = torch.Generator(device="cpu").manual_seed(model_seed + 7777) n_in = self.sizes[0] - # projection of output error back onto the input (fixed random) - self.Fproj = (torch.randn(n_in, self.n_classes, generator=g) - * (f_scale / math.sqrt(self.n_classes))).to(self.W[0].device, self.dtype) + if original_init: + for l, w in enumerate(self.W): + limit = math.sqrt(6.0 / self.sizes[l]) + w.copy_((torch.rand(w.shape, generator=g) * 2.0 * limit - limit) + .to(w.device, w.dtype)) + limit = math.sqrt(6.0 / n_in) + self.Fproj = ((torch.rand(n_in, self.n_classes, generator=g) * 2.0 * limit - limit) + * f_scale).to(self.W[0].device, self.dtype) + self.use_bias = use_bias + self.keep_prob = keep_prob + + def _pepita_forward(self, x, masks): + h = [x] + for l in range(self.L): + cur = h[-1] @ self.W[l].t() + if self.use_bias: + cur = cur + self.b[l] + if l < self.L - 1: + cur = self.act(cur) + if masks is not None: + cur = cur * masks[l] + h.append(cur) + return h def pepita_step(self, x, y, yoh, eta, momentum=0.0): with torch.no_grad(): - std = self.forward(x) - hs = std["h"] + masks = None + if self.keep_prob < 1.0: + masks = [(torch.rand(x.shape[0], self.sizes[l + 1], device=x.device) + < self.keep_prob).to(x.dtype) / self.keep_prob + for l in range(self.L - 1)] + hs = self._pepita_forward(x, masks) e = torch.softmax(hs[-1], 1) - yoh loss = F.cross_entropy(hs[-1], y).item() x_mod = x + (e @ self.Fproj.t()) # modulate the input by the error - mod = self.forward(x_mod) - hm = mod["h"] + hm = self._pepita_forward(x_mod, masks) + mod_error = torch.softmax(hm[-1], 1) - yoh B = x.shape[0] for l in range(self.L): - # change in this layer's output between clean and modulated passes - dpost = hs[l + 1] - hm[l + 1] - # PEPITA: first layer uses the CLEAN input as presynaptic; deeper - # layers use the modulated presynaptic. - pre = x if l == 0 else hm[l] - dW = -(dpost.t() @ pre) / B - db = -dpost.mean(0) + pre = hm[l] # modulated presynaptic state + if l == self.L - 1: + update = -(mod_error.t() @ pre) / B + update_b = -mod_error.mean(0) + else: + diff = hs[l + 1] - hm[l + 1] + update = -(diff.t() @ pre) / B + update_b = -diff.mean(0) if momentum: - self.mW[l].mul_(momentum).add_(dW); self.mb[l].mul_(momentum).add_(db) - self.W[l] += eta * self.mW[l]; self.b[l] += eta * self.mb[l] + self.mW[l].mul_(momentum).add_(update) + self.W[l] += eta * self.mW[l] + if self.use_bias: + self.mb[l].mul_(momentum).add_(update_b) + self.b[l] += eta * self.mb[l] else: - self.W[l] += eta * dW; self.b[l] += eta * db + self.W[l] += eta * update + if self.use_bias: + self.b[l] += eta * update_b return loss + def forward(self, x, **_): + h = self._pepita_forward(x, masks=None) + return {"h": h, "u": []} + # -------------------------------------------------------------------------- # Forward-Forward (Hinton 2022) @@ -107,8 +165,9 @@ class FFNet: each layer is trained by a local goodness objective; classification sums goodness across layers over the candidate labels.""" - def __init__(self, sizes, act="tanh", device="cpu", seed=0, threshold=2.0, - n_classes=10, dtype=torch.float32, overlay_val=10.0): + def __init__(self, sizes, act="relu", device="cpu", seed=0, threshold=2.0, + n_classes=10, dtype=torch.float32, overlay_val=None, + score_from_layer=1): self.sizes = list(sizes) self.L = len(sizes) - 1 self.n_classes = n_classes @@ -117,18 +176,19 @@ class FFNet: self.act, _ = ACTS[act] self.thr = threshold self.overlay_val = overlay_val + self.score_from_layer = min(score_from_layer, max(0, self.L - 1)) g = torch.Generator(device="cpu").manual_seed(seed) - self.W, self.b, self.mW, self.mb = [], [], [], [] + self.W, self.b, self.optimizers = [], [], [] for i in range(self.L): - w = (torch.randn(sizes[i + 1], sizes[i], generator=g) / math.sqrt(sizes[i]) - ).to(device, dtype) - w.requires_grad_(True) + # Match torch.nn.Linear initialization used by the public reference. + bound = 1.0 / math.sqrt(sizes[i]) + w = (torch.rand(sizes[i + 1], sizes[i], generator=g) * 2.0 * bound - bound).to( + device, dtype).requires_grad_(True) self.W.append(w) - bb = torch.zeros(sizes[i + 1], device=device, dtype=dtype, requires_grad=True) + bb = (torch.rand(sizes[i + 1], generator=g) * 2.0 * bound - bound).to( + device, dtype).requires_grad_(True) self.b.append(bb) - - def __init_overlay_scale__(self): - pass + self.optimizers.append(torch.optim.Adam([w, bb], lr=0.03)) def _overlay(self, x, labels): """Overlay one-hot label onto the first n_classes input features. @@ -136,7 +196,8 @@ class FFNet: (10 of 784 pixels is otherwise swamped by the shared image).""" xo = x.clone() xo[:, :self.n_classes] = 0.0 - xo[torch.arange(x.shape[0]), labels] = self.overlay_val + value = x.max().detach() if self.overlay_val is None else self.overlay_val + xo[torch.arange(x.shape[0], device=x.device), labels] = value return xo @staticmethod @@ -144,28 +205,40 @@ class FFNet: return h / (h.norm(dim=1, keepdim=True) + 1e-8) def _layer_forward(self, l, h_in): - return self.act(h_in @ self.W[l].t() + self.b[l]) + return self.act(self._norm(h_in) @ self.W[l].t() + self.b[l]) + + def _inputs_to_layer(self, l, x): + h = x + with torch.no_grad(): + for k in range(l): + h = self._layer_forward(k, h) + return h.detach() + + def train_layer(self, l, x, y, eta=0.03, negative_labels=None): + """One minibatch update of one greedy FF layer (local autograd only).""" + if negative_labels is None: + negative_labels = ((y + torch.randint( + 1, self.n_classes, y.shape, device=y.device)) % self.n_classes) + hpos = self._inputs_to_layer(l, self._overlay(x, y)) + hneg = self._inputs_to_layer(l, self._overlay(x, negative_labels)) + hp = self._layer_forward(l, hpos) + hn = self._layer_forward(l, hneg) + gp = hp.pow(2).mean(1) + gn = hn.pow(2).mean(1) + loss = (F.softplus(-gp + self.thr) + F.softplus(gn - self.thr)).mean() + opt = self.optimizers[l] + opt.param_groups[0]["lr"] = eta + opt.zero_grad() + loss.backward() + opt.step() + return loss.item() def train_step(self, x, y, eta): - # positive = correct label; negative = a random wrong label + """Compatibility path; canonical runner trains greedily layer-by-layer.""" neg = (y + torch.randint(1, self.n_classes, y.shape, device=y.device)) % self.n_classes - xpos = self._overlay(x, y) - xneg = self._overlay(x, neg) - hpos, hneg = xpos, xneg total = 0.0 for l in range(self.L): - hp = self._layer_forward(l, hpos.detach()) - hn = self._layer_forward(l, hneg.detach()) - gp = hp.pow(2).mean(1) # goodness (positive) - gn = hn.pow(2).mean(1) # goodness (negative) - # push positive goodness above threshold, negative below - loss = (F.softplus(-(gp - self.thr)) + F.softplus(gn - self.thr)).mean() - gW, gb = torch.autograd.grad(loss, [self.W[l], self.b[l]]) - with torch.no_grad(): - self.W[l] -= eta * gW - self.b[l] -= eta * gb - total += loss.item() - hpos, hneg = self._norm(hp).detach(), self._norm(hn).detach() + total += self.train_layer(l, x, y, eta, neg) return total / self.L @torch.no_grad() @@ -176,9 +249,8 @@ class FFNet: good = torch.zeros(x.shape[0], device=x.device) for l in range(self.L): h = self._layer_forward(l, h) - if l > 0: # skip first layer for scoring (Hinton) + if l >= self.score_from_layer: # paper excludes first layer good = good + h.pow(2).mean(1) - h = self._norm(h) scores[:, c] = good return scores.argmax(1) @@ -201,15 +273,20 @@ class EPNet: Updates are local (products of adjacent-layer rho's) with no backprop.""" def __init__(self, sizes, device="cpu", seed=0, beta=0.5, dt=0.5, - T_free=20, T_nudge=8, dtype=torch.float32): + T_free=20, T_nudge=4, dtype=torch.float32, + random_beta_sign=True): self.sizes = list(sizes) self.L = len(sizes) - 1 # number of weight layers self.device = device self.dtype = dtype self.beta, self.dt, self.T_free, self.T_nudge = beta, dt, T_free, T_nudge + self.random_beta_sign = random_beta_sign g = torch.Generator(device="cpu").manual_seed(seed) - self.W = [(torch.randn(sizes[i + 1], sizes[i], generator=g) / math.sqrt(sizes[i]) - ).to(device, dtype) for i in range(self.L)] + self.W = [] + for i in range(self.L): + limit = math.sqrt(6.0 / (sizes[i] + sizes[i + 1])) + self.W.append((torch.rand(sizes[i + 1], sizes[i], generator=g) + * 2.0 * limit - limit).to(device, dtype)) self.b = [torch.zeros(sizes[i + 1], device=device, dtype=dtype) for i in range(self.L)] @staticmethod @@ -218,28 +295,28 @@ class EPNet: @staticmethod def rhop(s): - return ((s > 0) & (s < 1)).to(s.dtype) + # Theano's clip derivative used by the authors is active at the bounds. + return ((s >= 0) & (s <= 1)).to(s.dtype) def _settle(self, x, y=None, beta=0.0, s=None, T=20): rx = self.rho(x) if s is None: - # init in the active region (rhop=0 at the clamp boundaries would - # otherwise freeze the dynamics); a feedforward warm start settles fast. - s = [] - below = rx - for i in range(self.L): - below = (below @ self.W[i].t() + self.b[i]).clamp(0, 1) - s.append(below) + # The reference implementation starts persistent particles at zero. + s = [torch.zeros(x.shape[0], n, device=x.device, dtype=x.dtype) + for n in self.sizes[1:]] for _ in range(T): new = [] for k in range(self.L): below = rx if k == 0 else self.rho(s[k - 1]) - pre = below @ self.W[k].t() + self.b[k] + # -dE/ds for E = ||rho(s)||^2/2 - b*rho(s) + # - rho(s_below) W^T rho(s). + drive = -self.rho(s[k]) + below @ self.W[k].t() + self.b[k] if k < self.L - 1: # top-down from layer above - pre = pre + self.rho(s[k + 1]) @ self.W[k + 1] + drive = drive + self.rho(s[k + 1]) @ self.W[k + 1] if k == self.L - 1 and beta: # nudge output toward target - pre = pre + beta * (y - s[k]) - ds = self.rhop(s[k]) * pre - s[k] + # Original cost is ||s_out-y||^2, hence the factor two. + drive = drive + 2.0 * beta * (y - s[k]) + ds = self.rhop(s[k]) * drive new.append((s[k] + self.dt * ds).clamp(0, 1)) s = new return s @@ -247,15 +324,19 @@ class EPNet: def train_step(self, x, y, yoh, eta): with torch.no_grad(): s0 = self._settle(x, beta=0.0, T=self.T_free) # free phase - sb = self._settle(x, yoh, beta=self.beta, s=[t.clone() for t in s0], T=self.T_nudge) + beta = self.beta + if self.random_beta_sign and torch.randint(0, 2, ()).item() == 0: + beta = -beta + sb = self._settle(x, yoh, beta=beta, s=[t.clone() for t in s0], T=self.T_nudge) B = x.shape[0] for k in range(self.L): below0 = self.rho(x) if k == 0 else self.rho(s0[k - 1]) belowb = self.rho(x) if k == 0 else self.rho(sb[k - 1]) - dW = (self.rho(sb[k]).t() @ belowb - self.rho(s0[k]).t() @ below0) / (self.beta * B) - db = (self.rho(sb[k]) - self.rho(s0[k])).mean(0) / self.beta - self.W[k] += eta * dW - self.b[k] += eta * db + dW = (self.rho(sb[k]).t() @ belowb - self.rho(s0[k]).t() @ below0) / (beta * B) + db = (self.rho(sb[k]) - self.rho(s0[k])).mean(0) / beta + layer_eta = eta[k] if isinstance(eta, (list, tuple)) else eta + self.W[k] += layer_eta * dW + self.b[k] += layer_eta * db # free-phase output as prediction proxy for loss logging return F.mse_loss(s0[-1], yoh).item() |
