AISTATS 2026

Finite-Time Analysis of Gradient Descent for Shallow Transformers

Enes Arda, Semih Cayci, Atilla Eryilmaz

Consider a shallow Transformer \(f(\cdot;\varphi)\colon\mathbb{R}^{d\times T}\to\mathbb{R}\) with \(m\) softmax attention heads, where \(T\) is the sequence length and \(d\) is the token dimension:

Shallow Transformer architecture: each of m heads applies softmax attention and a feed-forward layer; a weighted linear combination of the head outputs produces the scalar prediction f(X; φ).

Given \(n\) i.i.d. training samples \(\{(X^{(j)},y^{(j)})\}_{j=1}^{n}\), a standard training objective is to minimize the empirical mean-squared error

\[\hat{\mathcal{L}}_n(\varphi) =\frac{1}{n}\sum_{j=1}^{n} \bigl(f(X^{(j)};\varphi)-y^{(j)}\bigr)^2.\]

Under this objective, we study a very natural question:

How do width, number of samples and sequence length affect Transformer training?

Our main result makes this dependence explicit through a non-asymptotic bound:

Theorem 1 (informal)

Under the paper’s assumptions, after \(\tau\) steps of projected gradient descent with an appropriate step size and projection radius, the averaged parameters \(\bar\varphi^{(\tau)}\) satisfy, with high probability,

\[ \newcommand{\errorbound}[2]{ \underset{\substack{\strut\text{#2}\\\strut\text{error}}}{ \underbrace{#1\vphantom{\frac{d^2}{\sqrt{\tau}}}} \vphantom{\underbrace{\sqrt{\frac{d\log(n)}{m}}\vphantom{\frac{d^2}{\sqrt{\tau}}}}} } } \hat{\mathcal{L}}_n\!\left(\bar\varphi^{(\tau)}\right) \lesssim \errorbound{\frac{d^2}{\sqrt{\tau}}}{optimization} +\errorbound{\sqrt{\frac{d^3}{m}}}{linearization} +\errorbound{\sqrt{\frac{d\log(n)}{m}}}{approximation},\]

up to polylogarithmic factors in \(m\), with no hidden dependence on \(T\).

For a fixed training accuracy, the bound gives two takeaways:

  • Sequence-length independence. Training steps \(\tau\) and width \(m\) need not grow with \(T\).
  • Logarithmic width. Width logarithmic in sample size (\(m\gtrsim\log n\)) suffices.

Technical challenges and key ideas

Error decomposition in function space: optimization connects initialization to a finite-width network, linearization connects the network to its linear model, and approximation connects the linear model to the target in the NTK-RKHS ball. The optimization rate is proportional to τ to the power −1/2; the linearization and approximation rates are proportional to m to the power −1/2.
  • Optimization error. Softmax is key to sequence-length independence. Its Jacobian has a covariance structure that gives, for token norms at most one,
    \[\left\|X J_{\sigma_s} X^\top\right\|_{\mathrm{op}} \le 1\]
    uniformly in \(T\), which enables \(T\)-independent Lipschitz and smoothness constants. A Lyapunov argument for projected gradient descent then turns these into a \(T\)-independent optimization bound.
  • Approximation error. Working near initialization, we use transportation mappings to construct a linearized predictor for our target function \(f^\star\) inside a Transformer-NTK RKHS ball of radius \(\bar\nu\). Concentration over \(m\) heads, uniformly across \(n\) samples, gives the \(\bar\nu\sqrt{\log(n)/m}\) term, making logarithmic width sufficient for a fixed target class and accuracy. The factor \(\bar\nu\) also captures the effect of target complexity: a larger RKHS norm means a more complex function to approximate, so our bound requires more heads for the same guarantee.
  • Linearization error. Projection keeps parameters near initialization, where we can bound the Taylor remainder.

Comparison with recurrent networks

  • IndRNN. In a recurrent network, backpropagating through a sequence of length \(T\) multiplies per-step Jacobians \(T\) times. These products can grow exponentially in \(T\), so gradients explode and training becomes unstable.
  • Transformer. The softmax Jacobian has a covariance structure, so its norm is bounded independently of \(T\). Attention therefore avoids this source of gradient amplification, and training stays stable.

The trade-off is memory. Transformer retains the full input context, so its memory requirement grows with \(T\), while an IndRNN carries a fixed-size recurrent state. Our experiments illustrate this tradeoff: the Transformer trains stably as \(T\) grows, at the cost of memory.

Citation

BibTeX
@inproceedings{arda2026finite,
  title     = {Finite-Time Analysis of Gradient Descent for Shallow Transformers},
  author    = {Arda, Enes and Cayci, Semih and Eryilmaz, Atilla},
  booktitle = {Proceedings of The 29th International Conference on Artificial Intelligence and Statistics},
  pages     = {3988--3996},
  year      = {2026},
  editor    = {Khan, Emtiyaz and Li, Yingzhen and Solin, Arno and Ramdas, Aaditya},
  volume    = {300},
  series    = {Proceedings of Machine Learning Research},
  publisher = {PMLR},
  url       = {https://proceedings.mlr.press/v300/arda26a.html}
}