NTH

High-Dimensional Learning Dynamics of Attention-Indexed Models

AuthorsYizhou Xu, Margarita Sagitova, Lenka Zdeborová, Florent Krzakala

September 9, 2026 2 min read
Watch on YouTube
The one-line take

This work explains how different ways of parameterizing attention can determine whether large models learn useful structure or remain stuck, using a detailed high-dimensional theory.

Key results

400
SGD truncation experiment dimension

Dimension used to compare online SGD with finite-order moment truncations.

0.01
SGD learning rate

Learning rate in the softmax and MSE truncation experiment.

512
SGD batch size

Batch size in the numerical comparison between SGD and the deterministic theory.

800
Higher-dimensional illustration

Embedding dimension used in the untied-attention weak-recovery illustration.

What the paper found

This paper develops a high-dimensional theory for attention-indexed models, a framework broad enough to represent multi-layer, multi-head attention-only transformers and related tasks such as autoregressive prediction and matrix denoising. As embedding dimension d grows, the population-loss landscape collapses to finitely many trace order parameters, but online stochastic gradient descent generates an intrinsically infinite hierarchy of matrix moments; finite truncations nevertheless approximate the dynamics with exponentially decaying error. The main result is that attention parameterization creates an architectural implicit bias. Directly optimizing the attention matrix S can remain trapped on an uninformative manifold, whereas tied attention, S = WWᵀ, uses positive-semidefinite geometry to break symmetry and achieve weak recovery on the Θ(d² log d) sample scale under a nonzero cross-gradient condition. Untied attention, S = UVᵀ, behaves differently: pre-activation means evolve on a fast timescale, matrix moments on a slower one, and recovery on the same Θ(d² log d) scale occurs only if the fast phase selects a symmetry-breaking state. The theory is supported numerically with softmax and MSE settings at d = 400, learning rate 0.01, and batch size 512, while a higher-dimensional illustration uses d = 800. The results clarify why parameterizations used in modern systems such as ChatGPT-style transformer architectures can induce qualitatively different learning trajectories even when they represent similar predictors.

Original abstract

Attention mechanisms are central to modern foundation models, yet their training dynamics remain poorly understood, especially when the attention matrices have extensive rank. In this work, we study attention-indexed models, a broad framework that can represent multi-layer and multi-head attention architectures. First, we show that, in a suitable high-dimensional limit, the population-loss landscape is characterized by a finite set of trace order parameters. In contrast, online stochastic gradient descent (SGD) is governed by an infinite hierarchy of matrix moments, which we show can be exponentially well-approximated by a finite truncated system. Second, this framework reveals that attention parameterization itself can act as an architectural implicit bias. Direct optimization of an attention matrix $S\in\mathbb{R}^{d\times d}$ can remain trapped in an uninformative state. Tied attention ($S=WW^\top$) induces an automatic symmetry-breaking mechanism and yields weak recovery in $Θ(d^2\log d)$ samples. For untied attention, $S=UV^\top$, we uncover a fast-slow mechanism: the pre-activation mean first evolves on a fast timescale, while the overlaps evolve on a slower one. Weak recovery on the $Θ(d^2\log d)$ scale occurs when the state selected by the fast dynamics breaks the initial symmetry.

Read the original paper

More in Attention Mechanisms

Browse all 18 papers →
01Attention

CoWindow Attention: Full Causal Coverage Is a Collective Property

Jingze Shi, Zhangyang Peng, Xianduo Li, Yanlin Qi, Xiaotian Lin, Haoxian Chen, Liangdong Wang, Guang Liu, Yuyu Luo

CoWindow Attention makes long-context transformers faster by letting attention heads collectively cover the past instead of redundantly reading all of it.

Read analysis