ParaTTT: Parallelizing Test-Time Training Layers with Deep Memories
Abstract
Test-time training (TTT) layers replace the linear fast-weight memory of modern recurrent architectures with a small neural network trained online by gradient descent. This nonlinear memory is more expressive than the linear memories of Gated DeltaNet and Mamba-2, but is difficult to parallelize: existing implementations such as TTT-MLP rely on sequential, block-wise weight updates that run several times slower than chunkwise-parallel linear recurrences. Moreover, their block-wise execution only approximates the token-by-token online recurrence, whereas linear recurrences admit an exact chunkwise-parallel evaluation. We introduce ParaTTT (PTTT), an algorithm for parallelizing arbitrarily deep nonlinear memories by iteratively refining estimates of the inner-loop gradient trajectory. Starting from a guess of the hidden-layer gradients, each refinement consists of three chunkwise-parallel passes: a forward pass that evaluates the resulting feature trajectory, an exact gated-delta-rule pass that computes the readout trajectory and its feature gradients, and a backward pass that propagates these gradients into a new estimate. The fixed point is the exact online gradient descent recurrence, while decoding requires no refinement iterations. We prove geometric convergence for a two-layer nonlinear memory under sufficient contraction conditions. In practice, two iterations suffice. With two refinement steps at 4K context, PTTT evaluates the inner-loop trajectory 32× faster than TTT-MLP with the block approximation at matched error, while the full 340M model trains 4.9× faster. Models trained at 340M and 1.3B parameters match or exceed Gated DeltaNet and Mamba-2 and remain competitive with same-size Transformers on language-modeling perplexity and zero-shot accuracy, with better recall than the linear recurrences.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.