Text Classification
Transformers
JAX
pcmlp
feature-extraction
predictive-coding
local-loss
flax
tiny-model
custom-architecture
custom_code
Eval Results (legacy)
Instructions to use zeechimp/pc-mlp-tiny with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use zeechimp/pc-mlp-tiny with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="zeechimp/pc-mlp-tiny", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("zeechimp/pc-mlp-tiny", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download pc_mlp_tiny.py from zeechimp/pc-mlp-tiny: direct link, hf CLI and curl.
- Browser
- Download file 6.21 kB
-
https://huggingface.co/zeechimp/pc-mlp-tiny/resolve/main/pc_mlp_tiny.py
- Command line
-
hf download hf://zeechimp/pc-mlp-tiny/pc_mlp_tiny.py
-
curl -L -o pc_mlp_tiny.py https://huggingface.co/zeechimp/pc-mlp-tiny/resolve/main/pc_mlp_tiny.py
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) | |
| 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) |