pc-mlp-tiny / pc_mlp_tiny.py
zeechimp's picture
Upload pc_mlp_tiny.py
27c7eb0 verified
Raw History Blame Contribute Delete
6.21 kB
"""
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)