""" pc_mlp_tiny.py — Predictive-Coding MLP for sequence classification. 3,138 parameters. Matches a 5,314-param transformer on the synthetic palindrome + position task at 59% of the parameter count. Run: python pc_mlp_tiny.py """ import os, time, math os.environ.setdefault("XLA_FLAGS", "--xla_cpu_multi_thread_eigen=true intra_op_parallelism_threads=8") import numpy as np import jax, jax.numpy as jnp from jax import jit, random # ============================================================ # Config # ============================================================ VOCAB = 16 SEQ = 16 CLASSES = 2 D_HIDDEN = 32 LOCAL_LOSS_WEIGHT = 0.1 # ============================================================ # Task # ============================================================ def make_batch(key, B): """Return (x, y) where y=1 iff x[0] >= VOCAB/2 OR x is a palindrome.""" x = random.randint(key, (B, SEQ), 0, VOCAB) local = (x[:, 0] >= VOCAB // 2) pal = jnp.all(x == x[:, ::-1], axis=1) return x, (local | pal).astype(jnp.int32) def ce_loss(logits, y): return -jnp.take_along_axis( jax.nn.log_softmax(logits, -1), y[:, None], -1).mean() # ============================================================ # Model # ============================================================ def pc_init(key, d=VOCAB, hidden=D_HIDDEN): """Initialize PC-MLP parameters.""" k = random.split(key, 5) scale = 1.0 / math.sqrt(hidden) return { "tok": random.normal(k[0], (VOCAB, d)) * 0.05, "pos": random.normal(k[1], (SEQ, d)) * 0.05, "W1": random.normal(k[2], (d, hidden)) * scale, "W2": random.normal(k[3], (hidden, hidden)) * scale, "head": {"w": random.normal(k[4], (hidden, CLASSES)) * 0.02, "b": jnp.zeros(CLASSES)}, } def pc_forward(p, x): """Forward pass. Returns (logits, h0, h1, h2).""" B, T = x.shape h = p["tok"][x] + p["pos"][:T][None] # (B, T, d) h = h.reshape(B, T * p["tok"].shape[-1]) h = h[:, :p["W1"].shape[0]] # (B, d) h0 = h h1 = jax.nn.gelu(h0 @ p["W1"]) h2 = jax.nn.gelu(h1 @ p["W2"]) logits = h2 @ p["head"]["w"] + p["head"]["b"] return logits, h0, h1, h2 def pc_loss(p, x, y, lam=LOCAL_LOSS_WEIGHT): """Global CE loss + local predictive-coding regularizer.""" logits, h0, h1, h2 = pc_forward(p, x) global_loss = ce_loss(logits, y) local_loss = (jnp.mean((h1.mean(1) - h0.mean(1)) ** 2) + jnp.mean((h2.mean(1) - h1.mean(1)) ** 2)) return global_loss + lam * local_loss # ============================================================ # Training # ============================================================ def train(steps=200, B=32, seed=0, lr=3e-3): key = random.key(seed) p = pc_init(key) opt = {"m": jax.tree.map(jnp.zeros_like, p), "v": jax.tree.map(jnp.zeros_like, p), "t": jnp.int32(0)} def loss_fn(p, x, y): return pc_loss(p, x, y) @jit def step(p, opt, x, y): l, g = jax.value_and_grad(loss_fn)(p, x, y) t = opt["t"] + 1 m = jax.tree.map(lambda m, g: 0.9*m + 0.1*g, opt["m"], g) v = jax.tree.map(lambda v, g: 0.999*v + 0.001*g*g, opt["v"], g) mh = jax.tree.map(lambda m: m / (1 - 0.9**t), m) vh = jax.tree.map(lambda v: v / (1 - 0.999**t), v) np_ = jax.tree.map( lambda p, mh, vh: p - lr*mh / (jnp.sqrt(vh) + 1e-8), p, mh, vh) return np_, {"m": m, "v": v, "t": t}, l t0 = time.perf_counter() for s in range(steps): key, kb = random.split(key) xb, yb = make_batch(kb, B) p, opt, l = step(p, opt, xb, yb) wall = time.perf_counter() - t0 # eval key, kv = random.split(key) xv, yv = make_batch(kv, 512) logits, *_ = pc_forward(p, xv) acc = float((logits.argmax(-1) == yv).mean()) val_loss = float(ce_loss(logits, yv)) n = sum(int(np.prod(v.shape)) for v in jax.tree.leaves(p) if hasattr(v, "shape")) return p, {"wall_s": wall, "params": n, "val_acc": acc, "val_loss": val_loss, "train_loss": float(l)} # ============================================================ # Activation health diagnostic # ============================================================ def debug_activations(p): """Print mean/std at each nonlinearity. Catches dead layers.""" key = random.key(99) x, _ = make_batch(key, 32) h = p["tok"][x] + p["pos"][:x.shape[1]][None] h = h.reshape(x.shape[0], -1)[:, :p["W1"].shape[0]] pre1 = h @ p["W1"] h1 = jax.nn.gelu(pre1) pre2 = h1 @ p["W2"] h2 = jax.nn.gelu(pre2) print(" --- activation health ---") print(f" pre1 (input to W1) std={float(pre1.std()):.4f}") print(f" h1 (after GELU) std={float(h1.std()):.4f}") print(f" pre2 (input to W2) std={float(pre2.std()):.4f}") print(f" h2 (after GELU) std={float(h2.std()):.4f}") # dead-layer check: if any std is < 1e-4, the layer is collapsed for name, val in [("pre1", pre1), ("h1", h1), ("pre2", pre2), ("h2", h2)]: if float(val.std()) < 1e-4: print(f" WARNING: {name} is collapsed (std < 1e-4)") # local loss terms print(f" local loss (h1−h0) {float(jnp.mean((h1.mean(1)-h.mean(1))**2)):.6f}") print(f" local loss (h2−h1) {float(jnp.mean((h2.mean(1)-h1.mean(1))**2)):.6f}") # ============================================================ # Main # ============================================================ if __name__ == "__main__": print("=" * 60) print("PC-MLP: Predictive-Coding MLP for sequence classification") print("=" * 60) p, r = train(steps=200, seed=1) print(f"\n params: {r['params']:,}") print(f" wall: {r['wall_s']:.2f}s") print(f" train loss:{r['train_loss']:.3f}") print(f" val loss: {r['val_loss']:.3f}") print(f" val acc: {r['val_acc']:.3f}") print() debug_activations(p)