🤖 AI Summary
This study addresses the limitation that sequential gradient writes during test-time training hinder parallel scaling. By revealing the duality between forward evaluation and backpropagation in online gradient descent, this work proposes a parallel scan algorithm based on costate prediction. Through the introduction of a causal auxiliary network and a consistency loss function, it enables exact parallel computation of forward inference and backpropagation alongside weight-efficient updates. The proposed mechanism strictly recovers the performance of sequential online learners. Furthermore, during deployment, simply removing the auxiliary network natively supports token-by-token model updates, offering a novel paradigm for efficient, parallelized test-time training.
📝 Abstract
Test-time training (TTT) allows a model to improve its predictions at inference time by updating weights after every observed token. However, sequential gra- dient writes make parallel training difficult. We observe that, given layer inputs and activation gradients (costates), online gradient descent admits exact parallel scans for both forward evaluation and reverse backpropagation. GradLev lever- ages this duality: a causal auxiliary network predicts costates across all tokens in parallel; associative scans compute the adapted weights and forward activations and propagate gradients backward; and the resulting gradient targets supervise the predictor via a consistency loss. Exact consistency guarantees exact recovery of the sequential online learner. At deployment, the auxiliary predictor is discarded, and the model updates natively via token-by-token forward and backward passes.