summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/baseline_run.py241
-rw-r--r--experiments/baseline_smoke.py107
-rwxr-xr-xexperiments/baseline_sweep.sh37
-rw-r--r--sdil/local_baselines.py241
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()