MesaNet: Sequence Modeling by Locally Optimal Test-Time Training

Fuente: arXiv
Saved in:
Bibliographic Details
Main Authors: von Oswald, Johannes, Scherrer, Nino, Kobayashi, Seijin, Versari, Luca, Yang, Songlin, Schlegel, Maximilian, Maile, Kaitlin, Schimpf, Yanick, Sieberling, Oliver, Meulemans, Alexander, Saurous, Rif A., Lajoie, Guillaume, Frenkel, Charlotte, Pascanu, Razvan, Arcas, Blaise Agüera y, Sacramento, João
Format: Preprint
Published: 2025
Subjects:
Online Access:
Tags: Add Tag
No Tags, Be the first to tag this record!
_version_ 1866918046445076480
author von Oswald, Johannes
Scherrer, Nino
Kobayashi, Seijin
Versari, Luca
Yang, Songlin
Schlegel, Maximilian
Maile, Kaitlin
Schimpf, Yanick
Sieberling, Oliver
Meulemans, Alexander
Saurous, Rif A.
Lajoie, Guillaume
Frenkel, Charlotte
Pascanu, Razvan
Arcas, Blaise Agüera y
Sacramento, João
author_facet von Oswald, Johannes
Scherrer, Nino
Kobayashi, Seijin
Versari, Luca
Yang, Songlin
Schlegel, Maximilian
Maile, Kaitlin
Schimpf, Yanick
Sieberling, Oliver
Meulemans, Alexander
Saurous, Rif A.
Lajoie, Guillaume
Frenkel, Charlotte
Pascanu, Razvan
Arcas, Blaise Agüera y
Sacramento, João
contents Sequence modeling is currently dominated by causal transformer architectures that use softmax self-attention. Although widely adopted, transformers require scaling memory and compute linearly during inference. A recent stream of work linearized the softmax operation, resulting in powerful recurrent neural network (RNN) models with constant memory and compute costs such as DeltaNet, Mamba or xLSTM. These models can be unified by noting that their recurrent layer dynamics can all be derived from an in-context regression objective, approximately optimized through an online learning rule. Here, we join this line of work and introduce a numerically stable, chunkwise parallelizable version of the recently proposed Mesa layer (von Oswald et al., 2024), and study it in language modeling at the billion-parameter scale. This layer again stems from an in-context loss, but which is now minimized to optimality at every time point using a fast conjugate gradient solver. Through an extensive suite of experiments, we show that optimal test-time training enables reaching lower language modeling perplexity and higher downstream benchmark performance than previous RNNs, especially on tasks requiring long context understanding. This performance gain comes at the cost of additional flops spent during inference time. Our results are therefore intriguingly related to recent trends of increasing test-time compute to improve performance -- here by spending compute to solve sequential optimization problems within the neural network itself.
format Preprint
id arxiv_https___arxiv_org_abs_2506_05233
institution arXiv
publishDate 2025
record_format arxiv
spellingShingle MesaNet: Sequence Modeling by Locally Optimal Test-Time Training
von Oswald, Johannes
Scherrer, Nino
Kobayashi, Seijin
Versari, Luca
Yang, Songlin
Schlegel, Maximilian
Maile, Kaitlin
Schimpf, Yanick
Sieberling, Oliver
Meulemans, Alexander
Saurous, Rif A.
Lajoie, Guillaume
Frenkel, Charlotte
Pascanu, Razvan
Arcas, Blaise Agüera y
Sacramento, João
Machine Learning
Artificial Intelligence
Computation and Language
Sequence modeling is currently dominated by causal transformer architectures that use softmax self-attention. Although widely adopted, transformers require scaling memory and compute linearly during inference. A recent stream of work linearized the softmax operation, resulting in powerful recurrent neural network (RNN) models with constant memory and compute costs such as DeltaNet, Mamba or xLSTM. These models can be unified by noting that their recurrent layer dynamics can all be derived from an in-context regression objective, approximately optimized through an online learning rule. Here, we join this line of work and introduce a numerically stable, chunkwise parallelizable version of the recently proposed Mesa layer (von Oswald et al., 2024), and study it in language modeling at the billion-parameter scale. This layer again stems from an in-context loss, but which is now minimized to optimality at every time point using a fast conjugate gradient solver. Through an extensive suite of experiments, we show that optimal test-time training enables reaching lower language modeling perplexity and higher downstream benchmark performance than previous RNNs, especially on tasks requiring long context understanding. This performance gain comes at the cost of additional flops spent during inference time. Our results are therefore intriguingly related to recent trends of increasing test-time compute to improve performance -- here by spending compute to solve sequential optimization problems within the neural network itself.
title MesaNet: Sequence Modeling by Locally Optimal Test-Time Training
topic Machine Learning
Artificial Intelligence
Computation and Language
url https://arxiv.org/abs/2506.05233