Learning to (Learn at Test Time): RNNs with Expressive Hidden States

Yu SunXinhao LiKaran DalalJiarui XuArjun VikramGenghan ZhangYann DuboisXinlei ChenXiaolong WangSanmi Koyejo

article2025ICML371 citations

Introduces Test-Time Training layers that treat RNN hidden states as internal machine learning models updated via self-supervised gradient steps, achieving linear-time sequence modeling that continues to improve across long contexts where existing architectures plateau.

Listen

Modern language models face a fundamental trade-off between computational efficiency and the ability to process long context. Self-attention mechanisms achieve superior language understanding across long contexts, but their processing cost grows quadratically with sequence length. Conversely, recurrent neural networks (RNNs) offer linear time complexity and constant inference cost per token, yet their performance degrades across long sequences because fixed-size hidden states fail to capture complex dependencies among thousands or millions of tokens.

The article introduces and evaluates Test-Time Training (TTT) layers, a novel sequence modeling framework designed to retain linear computational complexity while significantly enhancing hidden state expressiveness. To overcome standard RNN memory limitations, the authors frame the hidden state itself as an internal machine learning model whose parameters are continuously updated on test sequences using self-supervised gradient descent.

The researchers implemented two concrete versions: TTT-Linear, which uses a linear model as its internal state, and TTT-MLP, which employs a two-layer multi-layer perceptron. They evaluated these architectures against competitive Transformer and Mamba (a modern RNN baseline) models at scales ranging from 125 million to 1.3 billion parameters across context windows up to 32,000 tokens using standard text benchmarks. To ensure practical execution on modern hardware, the authors implemented mini-batch test-time updates and a mathematically equivalent dual formulation that accelerates training more than fivefold on hardware accelerators.

The empirical findings demonstrate that TTT layers maintain performance comparable to leading baselines in short contexts while outperforming existing RNN architectures in long contexts. At an 8,000-token context length on the Pile dataset and a 32,000-token context length on the Books benchmark, both TTT variants achieved lower perplexity than Mamba. Furthermore, like Transformers, TTT layers steadily decrease prediction perplexity as context grows through 32,000 tokens, whereas Mamba's ability to utilize additional context plateaus after 16,000 tokens. Theoretical analyses also established that linear attention and standard self-attention represent special parametric and non-parametric cases within this broader TTT formulation.

These results indicate that embedding self-supervised learning directly into recurrent hidden states resolves the compression bottleneck of linear-time sequence models without incurring the prohibitive memory growth of full attention. For engineering deployments, TTT-Linear delivers constant inference latency per token and matches or beats standard training speeds, demonstrating immediate viability for long-sequence tasks. While TTT-MLP demonstrates higher theoretical expressiveness in long contexts, its hardware input/output complexity currently introduces notable wall-clock time overhead during generation.

Organizations developing large-scale language systems should evaluate TTT-Linear as a drop-in recurrent alternative for long-context workloads where full self-attention is cost-prohibitive. Prior to deploying deeper configurations like TTT-MLP in production, engineering teams should conduct targeted hardware and kernel optimization pilots to alleviate memory bandwidth bottlenecks. Future research should focus on extending evaluations to multi-million token sequences, exploring convolutional architectures for the internal learner, and building dedicated multi-device pipeline parallelization.

arXiv: 2407.04620
Cover for Learning to (Learn at Test Time): RNNs with Expressive Hidden States

Abstract

Self-attention performs well in long context but has quadratic complexity. Existing RNN layers have linear complexity, but their performance in long context is limited by the expressive power of their hidden states. We present a practical framework for instantiating sequence modeling layers with linear complexity and expressive hidden states. The key idea is to make the hidden state a machine learning model itself, and the update rule a step of self-supervised learning. Since the hidden state is updated by training even on test sequences, our layers are called Test-Time Training (TTT) layers. We consider two instantiations: TTT-Linear and TTT-MLP, whose hidden state is a linear model and a two-layer MLP respectively. We evaluate our instantiations at the scale of 125M to 1.3B parameters, comparing with a strong Transformer and Mamba, a modern RNN. Similar to Transformer, TTT-Linear and TTT-MLP can keep reducing perplexity by conditioning on more tokens, while Mamba cannot after 16k context. TTT-MLP still faces challenges in memory I/O, but shows larger potential in long context, pointing to a promising direction for future research.

Table of Contents

  • 1. Introduction
  • 2. Method
  • 2.1. TTT as updating a hidden state
  • 2.2. Training a network with TTT layers
  • 2.3. Learning a self-supervised task for TTT
  • 2.4. Parallelization with mini-batch TTT
  • 2.5. Dual form
  • 2.6. Theoretical equivalences
  • 2.7. Implementation details
  • 3. Experiments
  • 3.1. Short context: the Pile
  • 3.2. Long context: Books
  • 3.3. Wall-clock time
  • 4. Related Work
  • 4.1. Learning at Test Time
  • 4.1.1. TEST-TIME TRAINING
  • 4.1.2. FAST WEIGHTS
  • 4.2. Modern RNN layers
  • 4.3. Learning to Learn
  • 5. Future work
  • Impact statement
  • Acknowledgements
  • References
  • Appendix
  • A. Dual Form
  • A.1. Forward pass
  • A.2. Primal form
  • A.3. Dual form
  • A.4. Derivation
  • B. Nadaraya-Watson estimator
  • C. Experiment details

Knowls

  1. Knowl 1 — Test-Time Training Layer Framework

    model/method

    A Test-Time Training (TTT) layer is a sequence modeling layer that formulates context compression as self-supervised learning on test sequence tokens. The hidden state at time step tt is defined as the weight parameters WtW_t of an inner model f(⋅;W)f(\cdot; W), which maps input representations to output representations.

    Given an input sequence x1,…,xTx_1, \dots, x_T with xt∈Rdx_t \in \mathbb{R}^d, the layer defines three learnable linear projections parameterized in an outer loop: a training view projection matrix θK\theta_K, a label view projection matrix θV\theta_V, and a test view projection matrix θQ\theta_Q.

    At each time step tt, the inner model is trained on a self-supervised reconstruction task with inner loss function:

    ℓ(W;xt)=∥f(θKxt;W)−θVxt∥2\ell(W; x_t) = \| f(\theta_K x_t; W) - \theta_V x_t \|^2

    The hidden state is updated via a step of gradient descent:

    Wt=Wt−1−η∇Wℓ(Wt−1;xt)W_t = W_{t-1} - \eta \nabla_W \ell(W_{t-1}; x_t)

    where η\eta is the inner-loop learning rate.

    The layer output token ztz_t is produced by evaluating the updated inner model on the test view:

    zt=f(θQxt;Wt)z_t = f(\theta_Q x_t; W_t)

    Outer-loop parameters θ={θK,θV,θQ,… }\theta = \{\theta_K, \theta_V, \theta_Q, \dots\} and the rest of the network are optimized end-to-end using standard language modeling loss (such as next-token prediction) on training data by backpropagating through the inner-loop optimization trajectory.

  2. Knowl 2 — Mini-Batch Gradient Descent for TTT Layers

    model/method

    Sequential online gradient descent updates Wt=Wt−1−η∇ℓ(Wt−1;xt)W_t = W_{t-1} - \eta \nabla \ell(W_{t-1}; x_t) one token at a time, which cannot be parallelized across the sequence length. Full batch gradient descent computes all descent directions with respect to the initial state W0W_0 as Gt=∇ℓ(W0;xt)G_t = \nabla \ell(W_0; x_t), which allows full parallelization but restricts WtW_t to only a single gradient step away from W0W_0, constraining the effective optimization capacity.

    To balance hardware parallelism with multi-step optimization capability, TTT layers use mini-batch gradient descent with a mini-batch size bb (e.g., b=16b = 16). For sequence index t∈{1,…,T}t \in \{1, \dots, T\}, let t′=t−(t mod b)t' = t - (t \bmod b) denote the last time step of the preceding mini-batch (or 00 for the first mini-batch). The gradient step direction for token tt is computed using the frozen weights at the mini-batch boundary:

    Gt=∇Wℓ(Wt′;xt)G_t = \nabla_W \ell(W_{t'}; x_t)

    Within each mini-batch chunk of size bb, all bb gradient steps {Gt}t=t′+1t′+b\{G_t\}_{t=t'+1}^{t'+b} are computed in parallel. The hidden state at time step tt is then evaluated as a cumulative sum:

    Wt=W0−η∑s=1tGsW_t = W_0 - \eta \sum_{s=1}^t G_s

    This maintains an O(1)O(1) per-token recurrent update while enabling parallel execution across bb tokens on modern hardware accelerators.

  3. Knowl 3 — Dual Form Computation for Linear TTT Layers

    model/method

    In primal form, calculating and storing individual gradient matrices Gt∈Rd×dG_t \in \mathbb{R}^{d \times d} for each token requires bb outer products per mini-batch, creating severe memory I/O bottlenecks. The dual form mathematically produces the identical output representations Z=[z1,…,zb]Z = [z_1, \dots, z_b] and final mini-batch weights WbW_b using hardware-efficient matrix multiplications without explicitly materializing intermediate token-level weight matrices.

    For a linear inner model f(x)=Wxf(x) = W x with W0∈Rd×dW_0 \in \mathbb{R}^{d \times d} over an input mini-batch X=[x1,…,xb]∈Rd×bX = [x_1, \dots, x_b] \in \mathbb{R}^{d \times b} (with θK=θV=θQ=I\theta_K = \theta_V = \theta_Q = I and inner step size η\eta):

    The end-of-batch weights are updated via:

    Wb=W0−2η(W0X−X)XTW_b = W_0 - 2\eta (W_0 X - X) X^T

    The mini-batch output token matrix Z∈Rd×bZ \in \mathbb{R}^{d \times b} is computed via:

    Z=W0X−2ηΔZ = W_0 X - 2\eta \Delta

    where Δ∈Rd×b\Delta \in \mathbb{R}^{d \times b} is given by:

    Δ=(W0X−X)⋅mask(XTX)\Delta = (W_0 X - X) \cdot \text{mask}(X^T X)

    and mask(M)\text{mask}(M) denotes a strictly causal mask setting all upper-triangular entries (above the main diagonal) to zero.

    Computational Complexity: Within a mini-batch of size bb and feature dimension dd, the primal form requires O(b⋅d2)O(b \cdot d^2) operations with heavy memory access. The dual form requires O(b⋅d2)O(b \cdot d^2) operations for WbW_b and an additional O(b2⋅d)O(b^2 \cdot d) operations for ZZ. When b≪db \ll d (e.g., b=16,d≥768b=16, d \ge 768), the dual form runs significantly faster by utilizing accelerator Tensor Cores.

  4. Knowl 4 — Equivalence Between Linear Attention and Batch-GD TTT-Linear

    theoretical result

    Consider a TTT layer configured with:

    1. A linear inner model f(x)=Wxf(x) = W x,
    2. Batch gradient descent update rule evaluated at W0=0W_0 = 0 with inner learning rate η=1/2\eta = 1/2,
    3. Inner reconstruction loss ℓ(W;xt)=∥f(θKxt;W)−θVxt∥2\ell(W; x_t) = \| f(\theta_K x_t; W) - \theta_V x_t \|^2.

    Given any input sequence x1,…,xTx_1, \dots, x_T, the output rule zt=f(θQxt;Wt)z_t = f(\theta_Q x_t; W_t) produces the exact same output sequence z1,…,zTz_1, \dots, z_T as causal linear attention:

    zt=∑s=1t(θVxs)(θKxs)T(θQxt)z_t = \sum_{s=1}^t (\theta_V x_s) (\theta_K x_s)^T (\theta_Q x_t)

    Scope: This mathematical equivalence demonstrates that causal linear attention is a specific instantiation of a TTT layer using a linear hidden state, zero initialization, full-context batch gradient descent, and η=1/2\eta = 1/2.

  5. Knowl 5 — Equivalence Between Self-Attention and Nonparametric Nadaraya-Watson TTT

    theoretical result

    Consider a TTT layer where the inner model is a nonparametric Nadaraya-Watson kernel regression estimator f(x;x1,…,xt)f(x; x_1, \dots, x_t) operating over context tokens x1,…,xtx_1, \dots, x_t:

    f(x;x1,…,xt)=1∑s=1tκ(x,xs)∑s=1tκ(x,xs)ysf(x; x_1, \dots, x_t) = \frac{1}{\sum_{s=1}^t \kappa(x, x_s)} \sum_{s=1}^t \kappa(x, x_s) y_s

    with target labels ys=θVxsy_s = \theta_V x_s and asymmetric kernel function:

    κ(x,x′;θK,θQ)∝exp⁡((θKx)T(θQx′))\kappa(x, x'; \theta_K, \theta_Q) \propto \exp\left( (\theta_K x)^T (\theta_Q x') \right)

    parameterized by outer-loop matrices θK\theta_K and θQ\theta_Q.

    Given an input sequence x1,…,xTx_1, \dots, x_T, evaluating the prediction at test token x=xtx = x_t via zt=f(xt;x1,…,xt)z_t = f(x_t; x_1, \dots, x_t) yields the exact output sequence z1,…,zTz_1, \dots, z_T of standard softmax self-attention:

    zt=softmax((θQxt)T[θKx1,…,θKxt])[θVx1,…,θVxt]Tz_t = \text{softmax}\left( (\theta_Q x_t)^T [\theta_K x_1, \dots, \theta_K x_t] \right) [\theta_V x_1, \dots, \theta_V x_t]^T

    Scope: This establishes that standard self-attention is an instance of the TTT framework where the hidden state is a non-parametric model (the uncompressed Key-Value cache) whose memory grows as O(t)O(t).

  6. Knowl 6 — Architectural Instantiations: TTT-Linear and TTT-MLP

    model/method

    Two parametric instantiations of the TTT layer differ in their inner-loop model architecture ff:

    1. TTT-Linear: The inner model is a linear transformation with residual connection and Layer Normalization (LN):
    f(x)=x+LN(Wx)f(x) = x + \text{LN}(W x)

    where W∈Rd×dW \in \mathbb{R}^{d \times d} is the recurrent square hidden state.

    1. TTT-MLP: The inner model is a two-layer multi-layer perceptron with residual connection and Layer Normalization:
    f(x)=x+LN(MLP(x))f(x) = x + \text{LN}(\text{MLP}(x))

    where the intermediate hidden dimension is 4×4\times the input dimension dd, with a GELU non-linearity between layers.

    Both variants incorporate:

    • Learnable Initialization: Initial weights W0=θinitW_0 = \theta_{\text{init}} are learned as outer-loop parameters shared across sequences, which stabilizes optimization.
    • Input-Dependent Learning Rate: The inner step size η\eta is gated per token as:
    η(x)=ηbaseσ(θlr⋅x)\eta(x) = \eta_{\text{base}} \sigma(\theta_{\text{lr}} \cdot x)

    where σ\sigma is the sigmoid function, θlr\theta_{\text{lr}} is a learnable vector, and ηbase\eta_{\text{base}} is a fixed base learning rate scalar (1.01.0 for TTT-Linear and 0.10.1 for TTT-MLP).

    • Backbone Integration: TTT layers are embedded within a Mamba-style residual backbone containing 1D temporal convolutions and multiplicative gating before the TTT block.
  7. Knowl 7 — Dual Form Algorithm for Multilayer Inner Models

    algorithm

    The dual form algorithm computes the updated layer weights and output activations for a KK-layer MLP inner model over a mini-batch of bb tokens using only matrix multiplications, element-wise activations, and causal masking.

    Input: Initial layer weights W01,…,W0KW_0^1, \dots, W_0^K, mini-batch training views X^1=[θKx1,…,θKxb]\hat{X}^1 = [\theta_K x_1, \dots, \theta_K x_b], labels Y=[θVx1,…,θVxb]Y = [\theta_V x_1, \dots, \theta_V x_b], test views Xˉ1=[θQx1,…,θQxb]\bar{X}^1 = [\theta_Q x_1, \dots, \theta_Q x_b], activations σ1,…,σK\sigma_1, \dots, \sigma_K
    Output: Mini-batch end weights Wb1,…,WbKW_b^1, \dots, W_b^K, mini-batch output representations XˉK+1=[z1,…,zb]\bar{X}^{K+1} = [z_1, \dots, z_b]
    // Step 1: Initial forward pass on training views
    for k=1k = 1 to KK:
        Zk=W0kX^kZ^k = W_0^k \hat{X}^k
        X^k+1=σk(Zk)\hat{X}^{k+1} = \sigma_k(Z^k)
    // Step 2: Backward pass computing layer gradients
    ∇X^K+1ℓ=X^K+1−Y\nabla_{\hat{X}^{K+1}} \ell = \hat{X}^{K+1} - Y
    for k=Kk = K down to 1:
        ∇Zkℓ=σk′(Zk)⊙∇X^k+1ℓ\nabla_{Z^k} \ell = \sigma'_k(Z^k) \odot \nabla_{\hat{X}^{k+1}} \ell
        ∇X^kℓ=(W0k)T∇Zkℓ\nabla_{\hat{X}^k} \ell = (W_0^k)^T \nabla_{Z^k} \ell
        ∇W0kℓ=(∇Zkℓ)(X^k)T\nabla_{W_0^k} \ell = (\nabla_{Z^k} \ell) (\hat{X}^k)^T
        $W_b^k = W_0^k - \nabla_{W_0^k} \ell
    // Step 3: Dual forward pass on test views
    for k=1k = 1 to KK:
        Zˉk=W0kXˉk−(∇Zkℓ)⋅mask((X^k)TXˉk)\bar{Z}^k = W_0^k \bar{X}^k - (\nabla_{Z^k} \ell) \cdot \text{mask}((\hat{X}^k)^T \bar{X}^k)
        Xˉk+1=σk(Zˉk)\bar{X}^{k+1} = \sigma_k(\bar{Z}^k)
    return Wb1,…,WbK,XˉK+1W_b^1, \dots, W_b^K, \bar{X}^{K+1}

    This formulation avoids computing token-level outer products GtkG_t^k while exactly replicating the forward pass of the primal update rule.

  8. Knowl 8 — Ablation Study from Linear Attention to TTT-Linear

    data/table

    Stepwise ablations on the Pile dataset with context length T=2048T = 2048 for 125M-parameter models illustrate the impact of each component added to transform vanilla linear attention into TTT-Linear:

    Configuration Perplexity Difference
    Linear attention (Katharopoulos et al., 2020) 15.91 –
    Linear attn. improved 15.23 -0.68
    TTT equivalence 15.23 0.00
    + learnable W0W_0 15.27 +0.04
    + LN and residual in ff 14.05 -1.22
    + mini-batch TTT (b=16b=16) 12.35 -1.70
    + learnable η\eta 11.99 -0.36
    + Mamba backbone 11.09 -0.90

    The largest reduction in perplexity (-1.70) comes from switching from full-sequence batch gradient descent (b=2048b=2048) to mini-batch gradient descent (b=16b=16). Adding LayerNorm and a residual connection inside the inner model provides the second largest improvement (-1.22), followed by the integration of the temporal convolution Mamba backbone (-0.90).

  9. Knowl 9 — Long-Context Scaling of TTT Layers vs. Mamba and Transformers

    empirical result

    Evaluations across context lengths from 1k to 32k tokens on Books3 and the Pile for model scales between 125M and 1.3B parameters (under matched training FLOP budgets) demonstrate:

    1. Short Context (2k tokens): TTT-Linear, Mamba, and standard Transformers attain roughly matched validation perplexity.
    2. Long Context Scaling (8k to 32k tokens):
      • Mamba's perplexity reduction plateaus after 16k context, struggling to benefit from longer history.
      • TTT-Linear and TTT-MLP continue reducing perplexity monotonically out to 32k context, closely matching the scaling profile of full Transformers with long-context finetuning.
      • At 8k context on the Pile and 32k context on Books3, both TTT-Linear and TTT-MLP outperform Mamba.
    3. Expressiveness vs. Context Length: Under matched FLOP budgets, TTT-MLP performs worse than TTT-Linear at short context (1k–2k) due to parameter allocation overhead, but surpasses TTT-Linear at long context (16k–32k), indicating that a more expressive non-linear hidden state is better utilized over large context windows.
  10. Knowl 10 — Hardware and Memory Bottlenecks of TTT Layers

    limitation

    TTT layers face specific hardware efficiency and systems constraints:

    1. TTT-MLP Memory I/O Overhead: Although TTT-MLP is efficient in theoretical FLOPs, non-linear activation functions and normalization layers inside the inner model cannot be accelerated into single matrix multiplications by the dual form. Vector-Jacobian products (VJPs) for these operations incur significant memory I/O, resulting in substantially higher wall-clock inference and training latency relative to TTT-Linear.
    2. Memory Footprint Across Time: Retaining all intermediate hidden states W1,…,WTW_1, \dots, W_T for the outer-loop backward pass is memory-prohibitive. To train within GPU/TPU memory limits, gradient checkpointing through time must be used to save only the T/bT/b boundary states Wk⋅bW_{k \cdot b} across mini-batches, requiring activation recomputation during backward passes.

Coverage note — None was omitted; all primary architectural definitions, dual form algorithms, theoretical equivalence theorems, empirical scaling results, ablation data, and hardware limitations from the paper are covered.

References

  1. 1.Achiam, J., Adler, S., Agarwal, S., Ahmad, L., Akkaya, I., Aleman, F. L., Almeida, D., Altenschmidt, J., Altman, S., Anadkat, S., et al. Gpt-4 technical report. arXiv preprint arXiv:2303.08774, 2023.
  2. 2.Andrychowicz, M., Denil, M., Gomez, S., Hoffman, M. W., Pfau, D., Schaul, T., Shillingford, B., and De Freitas, N. Learning to learn by gradient descent by gradient descent. Advances in neural information processing systems, 29, 2016.
  3. 3.Beck, M., Pöppel, K., Spanring, M., Auer, A., Prudnikova, O., Kopp, M., Klambauer, G., Brandstetter, J., and Hochreiter, S. xlstm: Extended long short-term memory. arXiv preprint arXiv:2405.04517, 2024.
  4. 4.Bengio, Y., Bengio, S., and Cloutier, J. Learning a synaptic learning rule. Citeseer, 1990.
  5. 5.Bierens, H. J. The nadaraya-watson kernel regression function estimator. (Serie Research Memoranda; No. 1988-58). Faculty of Economics and Business Administration, Vrije Universiteit Amsterdam., 1988.
  6. 6.Bishop, C. M. and Nasrabadi, N. M. Pattern recognition and machine learning, volume 4. Springer, 2006.
  7. 7.Black, S., Biderman, S., Hallahan, E., Anthony, Q., Gao, L., Golding, L., He, H., Leahy, C., McDonell, K., Phang, J., et al. Gpt-neox-20b: An open-source autoregressive language model. arXiv preprint arXiv:2204.06745, 2022.
  8. 8.Bottou, L. and Vapnik, V. Local learning algorithms. Neural computation, 4(6):888–900, 1992.
  9. 9.Breiman, L., Meisel, W., and Purcell, E. Variable kernel estimates of multivariate densities. Technometrics, 19(2): 135–144, 1977.
  10. 10.Cai, Z. Weighted nadaraya–watson regression estimation. Statistics & probability letters, 51(3):307–318, 2001.
  11. 11.Chen, T., Xu, B., Zhang, C., and Guestrin, C. Training deep nets with sublinear memory cost, 2016.
  12. 12.Chen, Y.-C. A tutorial on kernel density estimation and recent advances. Biostatistics & Epidemiology, 1(1):161–187, 2017.
  13. 13.Clark, K., Guu, K., Chang, M.-W., Pasupat, P., Hinton, G., and Norouzi, M. Meta-learning fast weight language models. arXiv preprint arXiv:2212.02475, 2022.
  14. 14.Dao, T. and Gu, A. Transformers are ssms: Generalized models and efficient algorithms through structured state space duality. arXiv preprint arXiv:2405.21060, 2024.
  15. 15.De, S., Smith, S. L., Fernando, A., Botev, A., Cristian-Muraru, G., Gu, A., Haroun, R., Berrada, L., Chen, Y., Srinivasan, S., et al. Griffin: Mixing gated linear recurrences with local attention for efficient language models. arXiv preprint arXiv:2402.19427, 2024.
  16. 16.de Vries, H. In the long (context) run, 2023. URL https://www.harmdevries.com/post/context-length/. Accessed: 2024-06-24.
  17. 17.Finn, C., Abbeel, P., and Levine, S. Model-agnostic meta-learning for fast adaptation of deep networks. In International conference on machine learning, pp. 1126–1135. PMLR, 2017.
  18. 18.Gandelsman, Y., Sun, Y., Chen, X., and Efros, A. A. Test-time training with masked autoencoders. Advances in Neural Information Processing Systems, 2022.
  19. 19.Gao, L., Biderman, S., Black, S., Golding, L., Hoppe, T., Foster, C., Phang, J., He, H., Thite, A., Nabeshima, N., Presser, S., and Leahy, C. The pile: An 800gb dataset of diverse text for language modeling, 2020.
  20. 20.Geng, X. EasyLM: A Simple And Scalable Training Framework for Large Language Models. https://github.com/young-geng/EasyLM, mar 2023. https://github.com/young-geng/EasyLM.
  21. 21.Gu, A. and Dao, T. Mamba: Linear-time sequence modeling with selective state spaces. arXiv preprint arXiv:2312.00752, 2023.
  22. 22.Hardt, M. and Sun, Y. Test-time training on nearest neighbors for large language models. arXiv preprint arXiv:2305.18466, 2023.
  23. 23.Hendrycks, D. and Gimpel, K. Gaussian error linear units (gelus). arXiv preprint arXiv:1606.08415, 2016.
  24. 24.Hinton, G. E. and Plaut, D. C. Using fast weights to deblur old memories. In Proceedings of the ninth annual conference of the Cognitive Science Society, pp. 177–186, 1987.
  25. 25.Hochreiter, S. and Schmidhuber, J. Long short-term memory. Neural computation, 9(8):1735–1780, 1997.
  26. 26.Irie, K., Csordás, R., and Schmidhuber, J. The dual form of neural networks revisited: Connecting test time predictions to training patterns via spotlights of attention. In International Conference on Machine Learning, pp. 9639–9659. PMLR, 2022.
  27. 27.Kaplan, J., McCandlish, S., Henighan, T., Brown, T. B., Chess, B., Child, R., Gray, S., Radford, A., Wu, J., and Amodei, D. Scaling laws for neural language models. arXiv preprint arXiv:2001.08361, 2020.
  28. 28.Katharopoulos, A., Vyas, A., Pappas, N., and Fleuret, F. Transformers are rnns: Fast autoregressive transformers with linear attention. In International conference on machine learning, pp. 5156–5165. PMLR, 2020.
  29. 29.Kirsch, L. and Schmidhuber, J. Meta learning backpropagation and improving it. Advances in Neural Information Processing Systems, 34:14122–14134, 2021.
  30. 30.Kwon, W., Li, Z., Zhuang, S., Sheng, Y., Zheng, L., Yu, C. H., Gonzalez, J., Zhang, H., and Stoica, I. Efficient memory management for large language model serving with pagedattention. In Proceedings of the 29th Symposium on Operating Systems Principles, pp. 611–626, 2023.
  31. 31.Lake, B. M., Ullman, T. D., Tenenbaum, J. B., and Gershman, S. J. Building machines that learn and think like people. Behavioral and brain sciences, 40:e253, 2017.
  32. 32.Le, Q. V. Building high-level features using large scale unsupervised learning. In 2013 IEEE international conference on acoustics, speech and signal processing, pp. 8595–8598. IEEE, 2013.
  33. 33.Liu, H., Yan, W., Zaharia, M., and Abbeel, P. World model on million-length video and language with blockwise ringattention. arXiv preprint arXiv:2402.08268, 2024.
  34. 34.Maclaurin, D., Duvenaud, D., and Adams, R. Gradient-based hyperparameter optimization through reversible learning. In International conference on machine learning, pp. 2113–2122. PMLR, 2015.
  35. 35.Metz, L., Maheswaranathan, N., Cheung, B., and Sohl-Dickstein, J. Meta-learning update rules for unsupervised representation learning. arXiv preprint arXiv:1804.00222, 2018.
  36. 36.Peng, B., Goldstein, D., Anthony, Q., Albalak, A., Alcaide, E., Biderman, S., Cheah, E., Ferdinan, T., Hou, H., Kazienko, P., et al. Eagle and finch: Rwkv with matrix-valued states and dynamic recurrence. arXiv preprint arXiv:2404.05892, 2024.
  37. 37.Rosenblatt, F. The perceptron: a probabilistic model for information storage and organization in the brain. Psychological review, 65(6):386, 1958.
  38. 38.Schlag, I., Munkhdalai, T., and Schmidhuber, J. Learning associative inference using fast weight memory. arXiv preprint arXiv:2011.07831, 2020.
  39. 39.Schlag, I., Irie, K., and Schmidhuber, J. Linear transformers are secretly fast weight programmers. In International Conference on Machine Learning, pp. 9355–9366. PMLR, 2021.
  40. 40.Schmidhuber, J. Evolutionary principles in self-referential learning, or on learning how to learn: the meta-meta-... hook. PhD thesis, Technische Universität München, 1987.
  41. 41.Schmidhuber, J. Learning to control fast-weight memories: An alternative to dynamic recurrent networks. Neural Computation, 4(1):131–139, 1992.
  42. 42.Shazeer, N. Glu variants improve transformer, 2020.
  43. 43.Shleifer, S., Weston, J., and Ott, M. Normformer: Improved transformer pretraining with extra normalization. arXiv preprint arXiv:2110.09456, 2021.
  44. 44.Su, J., Lu, Y., Pan, S., Murtadha, A., Wen, B., and Liu, Y. Roformer: Enhanced transformer with rotary position embedding, 2023.
  45. 45.Sun, Y., Wang, X., Liu, Z., Miller, J., Efros, A., and Hardt, M. Test-time training with self-supervision for generalization under distribution shifts. In International Conference on Machine Learning, pp. 9229–9248. PMLR, 2020.
  46. 46.Thrun, S. and Pratt, L. Learning to learn: Introduction and overview. In Learning to learn, pp. 3–17. Springer, 1998.
  47. 47.Tieleman, T. and Hinton, G. Using fast weights to improve persistent contrastive divergence. In Proceedings of the 26th annual international conference on machine learning, pp. 1033–1040, 2009.
  48. 48.Touvron, H., Martin, L., Stone, K., Albert, P., Almahairi, A., Babaei, Y., Bashlykov, N., Batra, S., Bhargava, P., Bhosale, S., Bikel, D., Blecher, L., Ferrer, C. C., Chen, M., Cucurull, G., Esiobu, D., Fernandes, J., Fu, J., Fu, W., Fuller, B., Gao, C., Goswami, V., Goyal, N., Hartshorn, A., Hosseini, S., Hou, R., Inan, H., Kardas, M., Kerkez, V., Khabsa, M., Kloumann, I., Korenev, A., Koura, P. S., Lachaux, M.-A., Lavril, T., Lee, J., Liskovich, D., Lu, Y., Mao, Y., Martinet, X., Mihaylov, T., Mishra, P., Molybog, I., Nie, Y., Poulton, A., Reizenstein, J., Rungta, R., Saladi, K., Schelten, A., Silva, R., Smith, E. M., Subramanian, R., Tan, X. E., Tang, B., Taylor, R., Williams, A., Kuan, J. X., Xu, P., Yan, Z., Zarov, I., Zhang, Y., Fan, A., Kambadur, M., Narang, S., Rodriguez, A., Stojnic, R., Edunov, S., and Scialom, T. Llama 2: Open foundation and fine-tuned chat models, 2023.
  49. 49.Vincent, P., Larochelle, H., Bengio, Y., and Manzagol, P.-A. Extracting and composing robust features with denoising autoencoders. In ICML, pp. 1096–1103, 2008.
  50. 50.Wang, R., Sun, Y., Gandelsman, Y., Chen, X., Efros, A. A., and Wang, X. Test-time training on video streams. arXiv preprint arXiv:2307.05014, 2023.
  51. 51.Wolf, T., Debut, L., Sanh, V., Chaumond, J., Delangue, C., Moi, A., Cistac, P., Rault, T., Louf, R., Funtowicz, M., et al. Huggingface’s transformers: State-of-the-art natural language processing. arXiv preprint arXiv:1910.03771, 2019.
  52. 52.Xiong, W., Liu, J., Molybog, I., Zhang, H., Bhargava, P., Hou, R., Martin, L., Rungta, R., Sankararaman, K. A., Oguz, B., Khabsa, M., Fang, H., Mehdad, Y., Narang, S., Malik, K., Fan, A., Bhosale, S., Edunov, S., Lewis, M., Wang, S., and Ma, H. Effective long-context scaling of foundation models, 2023.
  53. 53.Yang, S., Wang, B., Shen, Y., Panda, R., and Kim, Y. Gated linear attention transformers with hardware-efficient training. arXiv preprint arXiv:2312.06635, 2023.
  54. 54.Yang, S., Wang, B., Zhang, Y., Shen, Y., and Kim, Y. Parallelizing linear transformers with the delta rule over sequence length. arXiv preprint arXiv:2406.06484, 2024.
  55. 55.Zhang, B. and Sennrich, R. Root mean square layer normalization, 2019.
  56. 56.Zhang, H., Berg, A. C., Maire, M., and Malik, J. Svm-knn: Discriminative nearest neighbor classification for visual category recognition. In 2006 IEEE Computer Society Conference on Computer Vision and Pattern Recognition (CVPR’06), volume 2, pp. 2126–2136. IEEE, 2006.

Citation

MLA
Sun, Y., et al. “Learning to (Learn at Test Time): RNNs with Expressive Hidden States”. arXiv, 2024, http://arxiv.org/abs/2407.04620v4.
APA
Sun, Y., Li, X., Dalal, K., Xu, J., Vikram, A., Zhang, G., Dubois, Y., Chen, X., Wang, X., Koyejo, S., Hashimoto, T., & Guestrin, C. (2024). Learning to (Learn at Test Time): RNNs with Expressive Hidden States. arXiv. http://arxiv.org/abs/2407.04620v4
Chicago
Sun, Y., X. Li, K. Dalal, et al. 2024. “Learning to (Learn at Test Time): RNNs with Expressive Hidden States”. arXiv. http://arxiv.org/abs/2407.04620v4.
Harvard
Sun, Y. et al. (2024) “Learning to (Learn at Test Time): RNNs with Expressive Hidden States”, arXiv [Preprint]. Available at: http://arxiv.org/abs/2407.04620v4.
Vancouver
1. Sun Y, Li X, Dalal K, et al (2024) Learning to (Learn at Test Time): RNNs with Expressive Hidden States. arXiv

BibTeX

@article{sun2024learning,
  title = {Learning to (Learn at Test Time): RNNs with Expressive Hidden States},
  author = {Sun, Yu and Li, Xinhao and Dalal, Karan and Xu, Jiarui and Vikram, Arjun and Zhang, Genghan and Dubois, Yann and Chen, Xinlei and Wang, Xiaolong and Koyejo, Sanmi and Hashimoto, Tatsunori and Guestrin, Carlos},
  year = {2024},
  journal = {arXiv},
  url = {http://arxiv.org/abs/2407.04620v4},
  eprint = {2407.04620}
}
Metadata:arXiv

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/