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)

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:

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

@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.
Downloads last month
18
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Evaluation results