Skip to content

Latest commit

 

History

History
49 lines (36 loc) · 2.05 KB

File metadata and controls

49 lines (36 loc) · 2.05 KB

Predictive-coding distillation

Documentation · Training · Research report

GPT-2 maps token and position embeddings through transformer blocks, a final layer norm, and the output projection. The predictive-coding wrapper adds an error tensor at each block boundary:

s_i = f_i(s_(i-1)) + error_i

Each training batch has two phases.

  1. Relax the error tensors while holding model weights fixed. The objective combines prediction energy with the task loss.
  2. Reconstruct the settled states and update model weights against detached local targets. The detach boundaries keep each block's weight gradient local.

Distillation uses the pretrained teacher's logits as the target. The homotopy parameter tau = error_lr × relaxation_steps controls how far the inner loop relaxes before each weight update. The schedule increases that budget through successive stages.

The inner loop uses torch.autograd.grad with error tensors as its leaves. The weight phase differentiates the local energy with respect to model parameters. Both phases run in PyTorch.

Continual learning

After distillation, the domain-training driver uses next-token cross-entropy. It supports predictive-coding updates and backpropagation through the same curriculum. The optional evidence gate scales settled errors per layer and feature during the PC weight update. Its statistics are accumulated across the batch and updated after each optimizer step.

See domain training for the commands and continual-learning results for the measured comparisons.

Inference

The learned tensors use GPT-2's standard state-dictionary structure. The MORK adapter converts these tensors and transformer operations into programs for the Rust engine. Full-forward inference emits the embedding, transformer, and output operations. Incremental decode uses a NumPy prompt prefill and MORK phases for subsequent tokens.

The code map links these components to their implementations.