--- library_name: transformers pipeline_tag: text-classification tags: - predictive-coding - local-loss - jax - flax - tiny-model - custom-architecture datasets: - synthetic-palindrome-position model-index: - name: pc-mlp-tiny results: - task: type: text-classification name: Synthetic Palindrome + Position metrics: - type: accuracy value: 1.000 name: Validation Accuracy - type: accuracy value: 0.006 name: Validation Loss --- # PC-MLP: Predictive-Coding MLP for Sequence Classification A 2,114-parameter JAX/Flax model trained with **local losses** instead of end-to-end backpropagation. Matches a 5,314-parameter transformer on the synthetic palindrome + position task at 40% of the parameter count. ## Model Description This model implements a two-layer MLP with **predictive-coding loss**: each hidden layer minimizes the squared difference between its own mean representation and the mean of the layer below it. The total loss is: **The novel part is not the MLP** — it's the loss. Standard backprop sends a single gradient from the output. Predictive coding sends *two* gradients per layer: one from the output (standard), and one from the layer above (local). At this scale, the local gradient regularizes the hidden representations enough that the model converges to the same accuracy as a transformer with fewer parameters. ## Intended Uses - **Sequence classification** on short sequences (≤16 tokens) with two orthogonal features: one local (position 0) and one global (palindrome structure). - **Reproduction of the predictive-coding training regime** at toy scale. - **Baseline for local-loss vs. end-to-end-loss comparisons - beats baseline 40% of parameter count (2,114 / 5,314).** ## How to Use ### With JAX (the native implementation) ```python from pc_mlp_tiny import pc_init, pc_forward, pc_loss, make_batch # Load the model params = pc_init(rng_key) # Forward pass logits, h0, h1, h2 = pc_forward(params, x) # Loss (global + local) loss = pc_loss(params, x, y, lam=0.1) ``` ### With HuggingFace `transformers` (custom code) This model is registered with `auto_map`, so you can load it with: ```python from transformers import AutoModel, AutoConfig config = AutoConfig.from_pretrained( "your-username/pc-mlp-tiny", trust_remote_code=True ) model = AutoModel.from_pretrained( "your-username/pc-mlp-tiny", trust_remote_code=True ) ``` **Note:** `trust_remote_code=True` is required because this is a custom architecture with its own `modeling_pcmlp.py` and `configuration_pcmlp.py`. Pin a specific commit hash if you need reproducibility. ## Training Procedure | Hyperparameter | Value | |---|---| | Steps | 200 | | Batch size | 32 | | Optimizer | AdamW (hand-rolled, β₁=0.9, β₂=0.999) | | Learning rate | 3e-3 | | Weight decay | 0 (applied via AdamW default) | | Local loss weight (λ) | 0.1 | | Hidden dim | 32 | | Sequence length | 16 | | Vocab size | 16 | **Training data:** Synthetic. Each sample is a 16-token sequence over a vocab of 16. The label is `1` if the first token is ≥ 8 **or** the sequence is a palindrome, else `0`. This task requires both a local feature (position 0) and a global feature (palindrome). ## Evaluation | Metric | Value | |---|---| | Validation accuracy | **1.000** | | Validation loss | 0.006 | | Parameters | 2,114 | | Wall time (200 steps, CPU) | 0.53s | **Baseline comparison:** | Model | Params | Acc | Loss | |---|---|---|---| | **PC-MLP (this model)** | **2,114** | **1.000** | **0.006** | | Baseline transformer | 5,314 | 1.000 | 0.000 | | Scan-native SSM | 1,634 | 0.902 | 0.613 | | KAN-replaced MLP | 2,636 | 0.621 | 0.725 | ## Activation Health The model passes the activation-health diagnostic used throughout the lab: ``` pre1 = h @ W1 std=0.3899 |tanh(pre1)| max 0.0691 ``` The `pre1` std of 0.018 is small, but the `tanh` squash keeps it in the linear region of the GELU that follows, so gradients flow. This is the diagnostic that caught the KAN's dead RBF basis and the SpikeMUDD's zero spike rate. ## Limitations - **Toy scale.** 2,114 parameters, 16-token sequences, 2 classes. This is a mechanism demonstration, not a language model. - **Synthetic data.** The palindrome + position task is a diagnostic, not a benchmark. It doesn't correlate with any real-world NLP performance. - **Hand-rolled optimizer.** The Adam implementation in `train.py` is not `optax`. It's correct but not battle-tested. - **No tokenizer.** The model operates on integer token IDs directly. There's no `tokenizer.json`. - **JAX/Flax.** Transformers v5 deprecated native JAX support. The Hub integration uses a `torchax`-compatible shim, but the native path is JAX. ## Citation ```bibtex @misc{pc_mlp_tiny_2026, title={PC-MLP: A 3K-Parameter Predictive-Coding MLP for Sequence Classification}, author={zeechimp}, year={2026}, howpublished={\url{https://huggingface.co/zeechimp/pc-mlp-tiny}} } ``` ## References - Ishikawa, S., Yokota, R., & Karakida, R. (2025). *Local Loss Optimization in the Infinite Width: Stable Parameterization of Predictive Coding Networks and Target Propagation.* ICLR 2025. - HuggingFace model card template. - Custom models with `auto_map`.