Transformers Learn In-Context by Gradient Descent

Johannes von OswaldEyvind NiklassonEttore RandazzoJoão SacramentoAlexander MordvintsevAndrey ZhmoginovMax Vladymyrov

article2023ICML889 citations

Proves that transformer forward passes mechanistically implement gradient descent during in-context learning by establishing a direct equivalence between linear self-attention layers and gradient-based optimization steps.

Listen

Modern artificial intelligence heavily relies on Transformer architectures, largely due to their ability to adapt to new information provided in prompt sequences without modifying model parameters—a capability known as in-context learning. Despite widespread deployment, the internal mechanisms governing how Transformers achieve this adaptation have remained poorly understood. Understanding these mechanisms is critical for explaining model behavior, improving computational efficiency, and architecting better architectures. The article addresses this gap by evaluating whether in-context learning operates as an internal optimization process, demonstrating that Transformers implicitly perform gradient descent within their forward pass.

To evaluate this relationship, the article introduces a mathematical construction proving that a single linear self-attention layer can execute a step of standard gradient descent on a regression loss. The authors then test this empirically by training various Transformer models—ranging from single linear layers to deep, multi-layer networks with multi-layer perceptrons—on linear and non-linear regression tasks across thousands of synthetic datasets. They measure cosine similarity, prediction errors, parameter interpolation, and generalization performance against explicit gradient descent baselines under both standard and out-of-distribution conditions.

The investigation yields several key findings. First, trained single-layer linear self-attention models converge directly to the theoretical gradient descent weight construction, matching the predictions and behavior of one-step gradient descent on both normal and out-of-distribution validation data. Second, deeper multi-layer Transformers consistently outperform standard gradient descent; they discover an accelerated optimization algorithm, termed GD++, which applies an iterative data transformation that reduces the condition number of the data covariance matrix to speed up learning. Third, incorporating multi-layer perceptrons allows Transformers to solve non-linear regression tasks by performing gradient descent on deep data representations, behaving equivalently to kernelized regression. Fourth, when input and target values are provided as separate sequence tokens, initial softmax self-attention layers learn a copying mechanism that formats tokens into paired representations, enabling subsequent layers to execute gradient-based updates.

These findings establish that training a Transformer across tasks operates as meta-learning, instantiating an internal optimization algorithm inside the network's forward computations. This mechanistic understanding demystifies prompt-based learning and shows that standard architectural components, such as multi-layer perceptrons and attention heads, serve complementary roles in feature extraction, data organization, and optimization. Furthermore, the analysis reveals that standard features like softmax attention and layer normalization can slightly degrade linear regression precision relative to linear attention, though two-head configurations mitigate these discrepancies.

Based on these results, researchers and engineers should explore targeted architectural refinements, such as embedding declarative optimization nodes into self-attention blocks to reduce layer depth requirements. Organizations should also evaluate linear attention mechanisms for tasks requiring fast internal linear adaptation. Further research is necessary before deploying these concepts to large production systems, particularly to evaluate whether the gradient descent framing scales to massive language models, complex token structures, noisy training distributions, and classification objectives.

Confidence in these findings is high for the evaluated domain of regression tasks and controlled small-scale models. However, readers should note the boundary conditions: the empirical validation focuses on synthetic linear and non-linear regression, and full-scale autoregressive language models involve higher distributional complexity that may rely on additional emergent mechanisms.

Cover for Transformers Learn In-Context by Gradient Descent

Abstract

At present, the mechanisms of in-context learning in Transformers are not well understood and remain mostly an intuition. In this paper, we suggest that training Transformers on auto-regressive objectives is closely related to gradient-based meta-learning formulations. We start by providing a simple weight construction that shows the equivalence of data transformations induced by 1) a single linear self-attention layer and by 2) gradient-descent (GD) on a regression loss. Motivated by that construction, we show empirically that when training self-attention-only Transformers on simple regression tasks either the models learned by GD and Transformers show great similarity or, remarkably, the weights found by optimization match the construction. Thus we show how trained Transformers become mesa-optimizers i.e. learn models by gradient descent in their forward pass. This allows us, at least in the domain of regression problems, to mechanistically understand the inner workings of in-context learning in optimized Transformers. Building on this insight, we furthermore identify how Transformers surpass the performance of plain gradient descent by learning an iterative curvature correction and learn linear models on deep data representations to solve non-linear regression tasks. Finally, we discuss intriguing parallels to a mechanism identified to be crucial for in-context learning termed induction-head (Olsson et al., 2022) and show how it could be understood as a specific case of in-context learning by gradient descent learning within Transformers.

Table of Contents

  • 1. Introduction
  • 2. Linear self-attention can emulate gradient descent on a linear regression task
  • Data transformations induced by gradient descent
  • Transformations induced by gradient descent and a linear self-attention layer can be equivalent
  • 3. Trained Transformers do mimic gradient descent on linear regression tasks
  • One-step of gradient descent vs. a single trained self-attention layer
  • Multiple steps of gradient descent vs. multiple layers of self-attention
  • Transformers solve nonlinear regression tasks by gradient descent on deep data representations
  • 4. Do self-attention layers build regression tasks?
  • 5. Discussion
  • Acknowledgments
  • References
  • A. Appendix
  • A.1. Proposition 1
  • A.2. Comparing the out-of-distribution behavior of trained Transformers and GD
  • A.3. Linear mode connectivity between the weight construction of Prop 1 and trained Transformers
  • A.4. Visualizing the trained Transformer weights
  • A.5. Proof and discussion of Proposition 3
  • A.6. Dampening the self-attention layer
  • A.7. Sine wave regression
  • A.8. Proposition 2 and connections between gradient descent, kernelized regression and kernel smoothing
  • A.9. Linear vs. softmax self-attention as well LayerNorm Transformers
  • A.10. Details of curvature correction
  • A.11. Phase transitions
  • A.12. Experimental details

Knowls

  1. Knowl 1 — Equivalence of Single-Head Linear Self-Attention to One Step of Gradient Descent

    theoretical result

    A single linear self-attention (LSA) layer without softmax normalization can implement an exact step of gradient descent on a linear regression mean-squared-error loss within its forward computation.

    Let the dataset be D={(xi,yi)}i=1N\mathcal{D} = \{(x_i, y_i)\}_{i=1}^N with inputs xi∈RNxx_i \in \mathbb{R}^{N_x} and targets yi∈RNyy_i \in \mathbb{R}^{N_y}. The squared-error loss for a linear model y(x)=Wxy(x) = W x parameterized by W∈RNy×NxW \in \mathbb{R}^{N_y \times N_x} is:

    L(W)=12N∑i=1N∥Wxi−yi∥2L(W) = \frac{1}{2N} \sum_{i=1}^N \|W x_i - y_i\|^2

    Taking one gradient descent step from initial weight W0∈RNy×NxW_0 \in \mathbb{R}^{N_y \times N_x} with learning rate η\eta yields the weight update:

    ΔW=−η∇WL(W0)=−ηN∑i=1N(W0xi−yi)xiT\Delta W = -\eta \nabla_W L(W_0) = -\frac{\eta}{N} \sum_{i=1}^N (W_0 x_i - y_i) x_i^T

    Representing in-context training examples as concatenated tokens ej=(xj,yj)∈RNx+Nye_j = (x_j, y_j) \in \mathbb{R}^{N_x + N_y} for j=1,…,Nj = 1, \dots, N, and the query/test token as eN+1=(xtest,−W0xtest)e_{N+1} = (x_{\text{test}}, -W_0 x_{\text{test}}), a 1-head linear self-attention layer computes:

    ej←ej+PVKTqj=ej+PWV(∑i=1Nei⊗eiWKTWQ)eje_j \leftarrow e_j + P V K^T q_j = e_j + P W_V \left( \sum_{i=1}^N e_i \otimes e_i W_K^T W_Q \right) e_j

    By setting the parameter matrices (in block form with Ix∈RNx×NxI_x \in \mathbb{R}^{N_x \times N_x}, Iy∈RNy×NyI_y \in \mathbb{R}^{N_y \times N_y}) as:

    WK=WQ=(Ix000),WV=(00W0−Iy),P=ηNINx+NyW_K = W_Q = \begin{pmatrix} I_x & 0 \\ 0 & 0 \end{pmatrix}, \quad W_V = \begin{pmatrix} 0 & 0 \\ W_0 & -I_y \end{pmatrix}, \quad P = \frac{\eta}{N} I_{N_x + N_y}

    the resulting layer dynamics on each token ej=(xj,yj)e_j = (x_j, y_j) are:

    ej←(xjyj)+ηNI∑i=1N(0W0xi−yi)xiTxj=(xjyj−ΔWxj)e_j \leftarrow \begin{pmatrix} x_j \\ y_j \end{pmatrix} + \frac{\eta}{N} I \sum_{i=1}^N \begin{pmatrix} 0 \\ W_0 x_i - y_i \end{pmatrix} x_i^T x_j = \begin{pmatrix} x_j \\ y_j - \Delta W x_j \end{pmatrix}

    For the query token eN+1=(xtest,−W0xtest)e_{N+1} = (x_{\text{test}}, -W_0 x_{\text{test}}), the updated yy-component is −(W0+ΔW)xtest-(W_0 + \Delta W) x_{\text{test}}, which yields the post-gradient-descent prediction (W0+ΔW)xtest(W_0 + \Delta W) x_{\text{test}} after multiplication by −1-1.

  2. Knowl 2 — GD++: In-Context Gradient Descent with Iterative Curvature Correction

    model/method

    Multi-layer linear self-attention Transformers can implement an accelerated gradient descent variant termed GD++\text{GD}^{++}, which combines a gradient descent step on model weights with an input data transformation that acts as an empirical curvature correction.

    Let X=[x1,…,xN]∈RNx×NX = [x_1, \dots, x_N] \in \mathbb{R}^{N_x \times N} denote the in-context input data matrix. GD++\text{GD}^{++} transforms every input vector according to:

    xj←H(X)xj=(INx−γXXT)xjx_j \leftarrow H(X) x_j = (I_{N_x} - \gamma X X^T) x_j

    where γ∈R\gamma \in \mathbb{R} is a scalar preconditioning parameter. A single-head linear self-attention layer implements this joint data transformation and gradient descent update in a single forward step using the weight construction:

    WK=WQ=(Ix000),WV=(Ix0W−Iy),P=(−γIx00ηNIy)W_K = W_Q = \begin{pmatrix} I_x & 0 \\ 0 & 0 \end{pmatrix}, \quad W_V = \begin{pmatrix} I_x & 0 \\ W & -I_y \end{pmatrix}, \quad P = \begin{pmatrix} -\gamma I_x & 0 \\ 0 & \frac{\eta}{N} I_y \end{pmatrix}

    Under this parameterization, the token update is:

    (xjyj)←(xjyj)+(−γXXTxj−ΔWxj)\begin{pmatrix} x_j \\ y_j \end{pmatrix} \leftarrow \begin{pmatrix} x_j \\ y_j \end{pmatrix} + \begin{pmatrix} -\gamma X X^T x_j \\ -\Delta W x_j \end{pmatrix}

    The data transformation alters the regression loss Hessian from H=XXT=UΣUTH = X X^T = U \Sigma U^T (with sorted eigenvalues λ1≥⋯≥λn≥0\lambda_1 \ge \dots \ge \lambda_n \ge 0) to:

    H++=(I−γXXT)X((I−γXXT)X)T=U(Σ−2γΣ2+γ2Σ3)UTH^{++} = (I - \gamma X X^T) X ((I - \gamma X X^T) X)^T = U (\Sigma - 2\gamma \Sigma^2 + \gamma^2 \Sigma^3) U^T

    The modified eigenvalues are given by f(λ,γ)=λ−2γλ2+γ2λ3f(\lambda, \gamma) = \lambda - 2\gamma \lambda^2 + \gamma^2 \lambda^3. When γ\gamma is chosen appropriately, the condition number κ++=λ1++/λn++\kappa^{++} = \lambda_1^{++} / \lambda_n^{++} of the loss Hessian is decreased, accelerating gradient descent across subsequent layers.

  3. Knowl 3 — In-Context Kernel Regression via Transformer Blocks with Preceding MLPs

    theoretical result

    Preceding a linear self-attention layer with a token-wise Multi-Layer Perceptron (MLP) enables in-context learning of non-linear functions by executing gradient descent on deep representations, equivalent to kernelized least-squares regression.

    Let an MLP with a residual connection transform tokens ej=(xj,yj)e_j = (x_j, y_j) as:

    ej←ej+(m~(xj),0)=(m(xj),yj)e_j \leftarrow e_j + (\tilde{m}(x_j), 0) = (m(x_j), y_j)

    where m(x)=x+m~(x)∈RDm(x) = x + \tilde{m}(x) \in \mathbb{R}^{D} maps inputs to a deep representation while leaving targets yjy_j unaffected.

    A subsequent linear self-attention layer configured with the weight construction for gradient descent performs the update:

    ej←(m(xj)yj−ΔWm(xj))e_j \leftarrow \begin{pmatrix} m(x_j) \\ y_j - \Delta W m(x_j) \end{pmatrix}

    which descends the kernelized mean squared error loss:

    L(W)=12N∑i=1N∥Wm(xi)−yi∥2L(W) = \frac{1}{2N} \sum_{i=1}^N \|W m(x_i) - y_i\|^2

    For an initialized weight W0=0W_0 = 0 and query token etest=(m(xtest),0)e_{\text{test}} = (m(x_{\text{test}}), 0), the resulting test prediction after a single Transformer block is:

    y^=−η∇WL(0)m(xtest)=∑i=1Nyim(xi)Tm(xtest)=∑i=1Nyik(xi,xtest)\hat{y} = -\eta \nabla_W L(0) m(x_{\text{test}}) = \sum_{i=1}^N y_i m(x_i)^T m(x_{\text{test}}) = \sum_{i=1}^N y_i k(x_i, x_{\text{test}})

    where k(xi,xtest)=m(xi)Tm(xtest)k(x_i, x_{\text{test}}) = m(x_i)^T m(x_{\text{test}}) is the inner-product kernel induced by the MLP representation, matching Nadaraya-Watson nonparametric kernel regression.

  4. Knowl 4 — Token Merging and Task Construction via Self-Attention Copying

    theoretical result

    A single-head self-attention layer with positional encodings can transform an alternating sequence of isolated input and target tokens into merged tokens containing input-target pairs (xj,yj)(x_j, y_j), satisfying the precondition required for downstream gradient descent emulation.

    Let inputs xj∈RNxx_j \in \mathbb{R}^{N_x} and targets yj∈RNyy_j \in \mathbb{R}^{N_y} be presented as separate alternating tokens with concatenated one-hot positional encodings pk∈R2N+1p_k \in \mathbb{R}^{2N+1}:

    e2j=(xjp2j),e2j+1=((0,yj)p2j+1)e_{2j} = \begin{pmatrix} x_j \\ p_{2j} \end{pmatrix}, \quad e_{2j+1} = \begin{pmatrix} (0, y_j) \\ p_{2j+1} \end{pmatrix}

    where 00 is a zero vector of dimension Nx−NyN_x - N_y.

    Setting projection matrix P=IP = I, and defining:

    WV=(00Ix−Ix,off),WK=(000Ix),WQ=(000Ix,offT)W_V = \begin{pmatrix} 0 & 0 \\ I_x & -I_{x,\text{off}} \end{pmatrix}, \quad W_K = \begin{pmatrix} 0 & 0 \\ 0 & I_x \end{pmatrix}, \quad W_Q = \begin{pmatrix} 0 & 0 \\ 0 & I_{x,\text{off}}^T \end{pmatrix}

    where Ix,offI_{x,\text{off}} is the lower diagonal identity matrix of size NxN_x, ensures that the attention dot product selects KTWQe2j=p2j+1K^T W_Q e_{2j} = p_{2j+1}. The attention update replaces the positional encoding of token e2je_{2j} with the target data from the adjacent token e2j+1e_{2j+1}:

    e2j←(xjp2j)+PVKTWQ(xjp2j)=(xjyj)e_{2j} \leftarrow \begin{pmatrix} x_j \\ p_{2j} \end{pmatrix} + P V K^T W_Q \begin{pmatrix} x_j \\ p_{2j} \end{pmatrix} = \begin{pmatrix} x_j \\ y_j \end{pmatrix}

    This exact copying operation holds under both linear self-attention and standard softmax self-attention.

  5. Knowl 5 — Approximation of Linear Gradient Descent via Two-Head Softmax Self-Attention

    theoretical result

    While single-head softmax self-attention introduces a task-dependent additive offset that prevents exact emulation of gradient descent, a two-head softmax self-attention layer can cancel this offset to recover linear gradient descent dynamics.

    Applying a first-order Taylor expansion to single-head softmax attention over context keys KK and query qj=WQxjq_j = W_Q x_j gives:

    softmax(KTqj)i=exiTWKQxj∑kexkTWKQxj≈1+xiTWKQxj∑k(1+xkTWKQxj)∝KTqj+ϵ\text{softmax}(K^T q_j)_i = \frac{e^{x_i^T W_{KQ} x_j}}{\sum_k e^{x_k^T W_{KQ} x_j}} \approx \frac{1 + x_i^T W_{KQ} x_j}{\sum_k (1 + x_k^T W_{KQ} x_j)} \propto K^T q_j + \epsilon

    where WKQ=WKTWQW_{KQ} = W_K^T W_Q and ϵ\epsilon is an additive error proportional to the sum over all input tokens.

    A two-head softmax self-attention layer cancels this linear offset by setting projection and value products to have opposing signs while sharing identical value mappings:

    P1V1softmax(K1TW1,Qxj)+P2V2softmax(K2TW2,Qxj)≈PV(x1T(W1,KQ−W2,KQ)xj⋮xNT(W1,KQ−W2,KQ)xj)∝PVKTqjP_1 V_1 \text{softmax}(K_1^T W_{1,Q} x_j) + P_2 V_2 \text{softmax}(K_2^T W_{2,Q} x_j) \approx P V \begin{pmatrix} x_1^T (W_{1,KQ} - W_{2,KQ}) x_j \\ \vdots \\ x_N^T (W_{1,KQ} - W_{2,KQ}) x_j \end{pmatrix} \propto P V K^T q_j

    When (W1,KQ−W2,KQ)(W_{1,KQ} - W_{2,KQ}) is diagonal, the two-head softmax layer successfully recovers the linear gradient descent construction.

  6. Knowl 6 — Empirical Convergence and Linear Mode Connectivity of Single-Layer Linear Self-Attention to Gradient Descent

    empirical result

    Single-layer linear self-attention Transformers trained on linear regression tasks converge to the explicit gradient descent weight construction θGD\theta_{\text{GD}} and exhibit linear mode connectivity with the theoretical construction.

    When evaluated over Tval=104T_{\text{val}} = 10^4 validation tasks (N=Nx=10N = N_x = 10, Ny=1N_y = 1):

    1. The cosine similarity between the input sensitivity ∂y^θ(xtest)∂xtest\frac{\partial \hat{y}_\theta(x_{\text{test}})}{\partial x_{\text{test}}} of the trained Transformer and ∂y^θGD(xtest)∂xtest\frac{\partial \hat{y}_{\theta_{\text{GD}}}(x_{\text{test}})}{\partial x_{\text{test}}} of a 1-step gradient descent model converges to ≈1.0\approx 1.0.
    2. Direct inspection of trained weight products WKQ=WKTWQW_{KQ} = W_K^T W_Q and WPV=PWVW_{PV} = P W_V matches the theoretical block diagonal structures.
    3. Linear interpolation between trained weights θ\theta and constructed weights θGD\theta_{\text{GD}}, defined by θI=12(θ+θGD)\theta_I = \frac{1}{2}(\theta + \theta_{\text{GD}}) after correcting for scalar ambiguity β=mean(diag(WKQ))\beta = \text{mean}(\text{diag}(W_{KQ})), incurs zero loss penalty and produces identical prediction loss.
    4. When evaluated out-of-distribution by scaling input ranges x∼U(−α,α)Nxx \sim \mathcal{U}(-\alpha, \alpha)^{N_x} or teacher weights αW\alpha W with α∈[0.5,2.0]\alpha \in [0.5, 2.0], the trained Transformer's loss curve perfectly overlays that of 1-step gradient descent.
  7. Knowl 7 — Alignment of Deep and Recurrent Linear Self-Attention Transformers with GD++

    empirical result

    Trained multi-layer and recurrent linear self-attention Transformers outperform standard KK-step gradient descent and align with GD++\text{GD}^{++} (gradient descent with learned data preconditioning).

    Key empirical findings across architectures on linear regression benchmarks (N=Nx=10,Ny=1N = N_x = 10, N_y = 1):

    1. 2-Layer Recurrent Transformer: Applying the same trained LSA layer twice achieves a lower MSE loss (≈0.08\approx 0.08) than 2 steps of standard GD (≈0.12\approx 0.12). It achieves cosine similarity ≈1.0\approx 1.0 with GD++\text{GD}^{++} using meta-learned η\eta and γ\gamma.
    2. 5-Layer Deep Transformer (Non-recurrent): A 5-layer unshared LSA model surpasses 5-step standard GD and aligns with 5-step GD++\text{GD}^{++} having layer-specific (ηk,γk)(\eta_k, \gamma_k) pairs.
    3. 10-Layer Recurrent & 12-Layer MLP-Transformer: Models with 10 recurrent layers and 12-layer blocks containing 4-head attention and MLPs closely match KK-step GD++\text{GD}^{++}.
    4. Out-of-Distribution Generalization: When tested on input distributions not seen during training (Normal, Exponential, and Laplace distributions scaled by α\alpha), recurrent Transformers exhibit loss profiles matching GD++\text{GD}^{++}.
  8. Knowl 8 — Nonlinear In-Context Learning via MLP-Transformer Blocks Matching Meta-Learned MLPs

    empirical result

    Transformers equipped with MLP layers and linear self-attention solve non-linear in-context regression tasks by matching the behavior of meta-learned feature representations adapted by output-layer gradient descent.

    On sine wave regression tasks y=asin⁡(ρ+x)y = a \sin(\rho + x) where a∼U(0.1,5)a \sim \mathcal{U}(0.1, 5), ρ∼U(0,π)\rho \sim \mathcal{U}(0, \pi), x∼U(−5,5)x \sim \mathcal{U}(-5, 5), and N=10N = 10 context points:

    1. A Transformer block consisting of an input affine embedding (dimension 40), a single-hidden-layer MLP (160 hidden units with GELU activation), and a linear self-attention layer was trained end-to-end.
    2. A control model sharing the same embedding and MLP architecture was trained via meta-learning (backpropagation through training) with 1 step of gradient descent applied to an explicit linear readout layer.
    3. Both models produce nearly identical initial representations before the attention/GD step and converge to the same final loss (≈0.001\approx 0.001).
    4. Prediction differences and partial derivative differences between the trained Transformer and the meta-learned GD model remain close to zero throughout training, with cosine similarity between sensitivities exceeding 0.950.95.
  9. Knowl 9 — Emergence of Induction-Like Copying in Multi-Layer Attention for In-Context Regression

    empirical result

    When trained on separated input and target tokens (e2j=(xj)e_{2j} = (x_j), e2j+1=(0,yj)e_{2j+1} = (0, y_j)), a two-layer attention-only Transformer exhibits a discrete phase transition during meta-training where the first layer discovers a copying mechanism immediately prior to achieving gradient descent performance.

    Tracking the partial derivative norms of the first layer's output t(ej)t(e_j) with respect to input tokens during online training reveals:

    1. During initial training steps (0 to ≈20,000\approx 20{,}000), loss remains high (≈0.40\approx 0.40) and sensitivity to other tokens is near zero.
    2. Immediately before the loss drops to match 1-step gradient descent (≈0.20\approx 0.20), the norm of the partial derivative ∥∂t(ej)∂ej+1∥\left\| \frac{\partial t(e_j)}{\partial e_{j+1}} \right\| spikes from ≈0\approx 0 to >3.0> 3.0.
    3. Sensitivities to non-adjacent tokens ∥∂t(ej)∂eother∥\left\| \frac{\partial t(e_j)}{\partial e_{\text{other}}} \right\| remain low (<0.5< 0.5).
    4. This demonstrates that layer 1 learns to copy target data yjy_j into the token holding input xjx_j, after which layer 2 executes 1-step gradient descent on the constructed in-context pairs.
  10. Knowl 10 — Online Meta-Learning Benchmark Protocol for In-Context Regression

    experimental setup

    The standard evaluation protocol for analyzing in-context gradient descent across linear and nonlinear tasks consists of the following parameters and procedures:

    • Task Distribution: For each task τ\tau, teacher weights are sampled as Wτ∼N(0,INy×Nx)W_\tau \sim \mathcal{N}(0, I_{N_y \times N_x}) and inputs as xτ,i∼U(−1,1)Nxx_{\tau, i} \sim \mathcal{U}(-1, 1)^{N_x}. Clean targets are generated as yτ,i=Wτxτ,iy_{\tau, i} = W_\tau x_{\tau, i}. Unless specified otherwise, default dimensions are N=Nx=10N = N_x = 10 and Ny=1N_y = 1.
    • Online Meta-Optimization: Model parameters θ\theta are optimized over non-repeating task batches by minimizing the expected MSE loss on the query token:

    L(θ)=1B∑τ=1B∥y^θ({eτ,i}i=1N,eτ,N+1)−yτ,test∥2\mathcal{L}(\theta) = \frac{1}{B} \sum_{\tau=1}^B \|\hat{y}_\theta(\{e_{\tau, i}\}_{i=1}^N, e_{\tau, N+1}) - y_{\tau, \text{test}}\|^2

    • Optimization Hyperparameters: Adam optimizer (learning rate 0.0010.001 for depth K<3K < 3, 0.00050.0005 for K≥3K \ge 3), batch size B=2048B = 2048, global gradient norm clipping at 10.010.0 via Optax.
    • Weight Initialization: Haiku fan-in truncated normal initialization with standard deviation 0.002/K0.002 / K, where KK is the number of layers.
    • Stability Control: For deeper architectures (K>2K > 2), token activations are clipped to [−10,10][-10, 10] after each layer to prevent divergence near the optimal learning rate regime.

Coverage note — Omitted qualitative grokking observations on small fixed datasets (Section A.11) and speculative future extensions involving declarative nodes and HyperTransformers (Section 5) as they are secondary to the core theoretical and empirical proof of in-context gradient descent.

References

  1. 1.Akyürek, E., Schuurmans, D., Andreas, J., Ma, T., and Zhou, D. What learning algorithm is in-context learning? investigations with linear models. In The Eleventh International Conference on Learning Representations, 2023. URL https://openreview.net/forum?id=0g0X4H8yN4I.
  2. 2.Amos, B. and Kolter, J. Z. Optnet: Differentiable optimization as a layer in neural networks. In International Conference on Machine Learning, 2017.
  3. 3.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. In Advances in Neural Information Processing Systems, 2016.
  4. 4.Ba, J., Hinton, G. E., Mnih, V., Leibo, J. Z., and Ionescu, C. Using fast weights to attend to the recent past. In Advances in Neural Information Processing Systems 29, 2016.
  5. 5.Bai, S., Kolter, J. Z., and Koltun, V. Deep equilibrium models. Advances in Neural Information Processing Systems, 2019.
  6. 6.Bengio, Y., Bengio, S., and Cloutier, J. Learning a synaptic learning rule. Technical report, Université de Montréal, Département d’Informatique et de Recherche opérationnelle, 1990.
  7. 7.Benzing, F., Schug, S., Meier, R., von Oswald, J., Akram, Y., Zucchet, N., Aitchison, L., and Steger, A. Random initialisations performing above chance and how to find them. OPT2022: 14th Annual Workshop on Optimization for Machine Learning, 2022.
  8. 8.Bertinetto, L., Henriques, J. F., Torr, P. H. S., and Vedaldi, A. Meta-learning with differentiable closed-form solvers. In International Conference on Learning Representations, 2019.
  9. 9.Brown, T. B., Mann, B., Ryder, N., Subbiah, M., Kaplan, J., Dhariwal, P., Neelakantan, A., Shyam, P., Sastry, G., Askell, A., Agarwal, S., Herbert-Voss, A., Krueger, G., Henighan, T., Child, R., Ramesh, A., Ziegler, D. M., Wu, J., Winter, C., Hesse, C., Chen, M., Sigler, E., Litwin, M., Gray, S., Chess, B., Clark, J., Berner, C., McCandlish, S., Radford, A., Sutskever, I., and Amodei, D. Language models are few-shot learners. arXiv preprint arXiv:2005.14165, 2020.
  10. 10.Carion, N., Massa, F., Synnaeve, G., Usunier, N., Kirillov, A., and Zagoruyko, S. End-to-end object detection with transformers. In Computer Vision – ECCV 2020. Springer International Publishing, 2020.
  11. 11.Chalmers, D. J. The evolution of learning: an experiment in genetic connectionism. In Touretzky, D. S., Elman, J. L., Sejnowski, T. J., and Hinton, G. E. (eds.), Connectionist Models, pp. 81–90. Morgan Kaufmann, 1991.
  12. 12.Chan, S. C. Y., Dasgupta, I., Kim, J., Kumaran, D., Lampinen, A. K., and Hill, F. Transformers generalize differently from information stored in context vs in weights. arXiv preprint arXiv:2210.05675, 2022a.
  13. 13.Chan, S. C. Y., Santoro, A., Lampinen, A. K., Wang, J. X., Singh, A., Richemond, P. H., McClelland, J., and Hill, F. Data distributional properties drive emergent in-context learning in transformers. Advances in Neural Information Processing Systems, 2022b.
  14. 14.Choromanski, K. M., Likhosherstov, V., Dohan, D., Song, X., Gane, A., Sarlos, T., Hawkins, P., Davis, J. Q., Mohiuddin, A., Kaiser, L., Belanger, D. B., Colwell, L. J., and Weller, A. Rethinking attention with performers. In International Conference on Learning Representations, 2021. URL https://openreview.net/forum?id=Ua6zuk0WRH.
  15. 15.Dai, D., Sun, Y., Dong, L., Hao, Y., Ma, S., Sui, Z., and Wei, F. Why can GPT learn in-context? language models implicitly perform gradient descent as meta-optimizers. In ICLR 2023 Workshop on Mathematical and Empirical Understanding of Foundation Models, 2023. URL https://openreview.net/forum?id=fzbHRjAd8U.
  16. 16.Dosovitskiy, A., Beyer, L., Kolesnikov, A., Weissenborn, D., Zhai, X., Unterthiner, T., Dehghani, M., Minderer, M., Heigold, G., Gelly, S., Uszkoreit, J., and Houlsby, N. An image is worth 16x16 words: Transformers for image recognition at scale. In International Conference on Learning Representations, 2021. URL https://openreview.net/forum?id=YicbFdNTTy.
  17. 17.Entezari, R., Sedghi, H., Saukh, O., and Neyshabur, B. The role of permutation invariance in linear mode connectivity of neural networks. arXiv preprint arXiv:2110.06296, 2021.
  18. 18.Finn, C. and Levine, S. Meta-learning and universality: Deep representations and gradient descent can approximate any learning algorithm. In International Conference on Learning Representations, 2018. URL https://openreview.net/forum?id=HyjC5yWCW.
  19. 19.Finn, C., Abbeel, P., and Levine, S. Model-agnostic meta-learning for fast adaptation of deep networks. In International Conference on Machine Learning, 2017.
  20. 20.Flennerhag, S., Rusu, A. A., Pascanu, R., Visin, F., Yin, H., and Hadsell, R. Meta-learning with warped gradient descent. In International Conference on Learning Representations, 2020.
  21. 21.Garg, S., Tsipras, D., Liang, P., and Valiant, G. What can transformers learn in-context? a case study of simple function classes. In Oh, A. H., Agarwal, A., Belgrave, D., and Cho, K. (eds.), Advances in Neural Information Processing Systems, 2022. URL https://openreview.net/forum?id=flNZJ2eOet.
  22. 22.Gordon, J., Bronskill, J., Bauer, M., Nowozin, S., and Turner, R. Meta-learning probabilistic inference for prediction. In International Conference on Learning Representations, 2019. URL https://openreview.net/forum?id=HkxStoC5F7.
  23. 23.Gould, S., Hartley, R., and Campbell, D. J. Deep declarative networks. IEEE Transactions on Pattern Analysis and Machine Intelligence, 2021.
  24. 24.Gulati, A., Qin, J., Chiu, C.-C., Parmar, N., Zhang, Y., Yu, J., Han, W., Wang, S., Zhang, Z., Wu, Y., and Pang, R. Conformer: Convolution-augmented transformer for speech recognition. arXiv preprint arXiv:2005.08100, 2020.
  25. 25.Hendrycks, D. and Gimpel, K. Gaussian error linear units (gelus). arXiv preprint arXiv:1606.08415, 2016.
  26. 26.Hinton, G. E. and Plaut, D. C. Using fast weights to deblur old memories. 1987.
  27. 27.Hochreiter, S., Younger, A. S., and Conwell, P. R. Learning to learn using gradient descent. In Dorffner, G., Bischof, H., and Hornik, K. (eds.), Artificial Neural Networks — ICANN 2001, pp. 87–94, Berlin, Heidelberg, 2001. Springer Berlin Heidelberg. ISBN 978-3-540-44668-2.
  28. 28.Hubinger, E., van Merwijk, C., Mikulik, V., Skalse, J., and Garrabrant, S. Risks from learned optimization in advanced machine learning systems. arXiv [cs.AI], Jun 2019. URL http://arxiv.org/abs/1906.01820.
  29. 29.Irie, K., Schlag, I., Csordás, R., and Schmidhuber, J. Going beyond linear transformers with recurrent fast weight programmers. CoRR, abs/2106.06295, 2021. URL https://arxiv.org/abs/2106.06295.
  30. 30.Kingma, D. P. and Ba, J. Adam: A method for stochastic optimization, 2014.
  31. 31.Kirsch, L. and Schmidhuber, J. Meta learning backpropagation and improving it. In Beygelzimer, A., Dauphin, Y., Liang, P., and Vaughan, J. W. (eds.), Advances in Neural Information Processing Systems, 2021. URL https://openreview.net/forum?id=hhU9TEvB6AF.
  32. 32.Kirsch, L., Harrison, J., Sohl-Dickstein, J., and Metz, L. General-purpose in-context learning by meta-learning transformers. In Sixth Workshop on Meta-Learning at the Conference on Neural Information Processing Systems, 2022. URL https://openreview.net/forum?id=t6tA-KB4dO.
  33. 33.Lee, K., Maji, S., Ravichandran, A., and Soatto, S. Meta-learning with differentiable convex optimization. In IEEE/CVF Conference on Computer Vision and Pattern Recognition, 2019.
  34. 34.Lee, Y. and Choi, S. Gradient-based meta-learning with learned layerwise metric and subspace. In International Conference on Machine Learning, 2018.
  35. 35.Li, Z., Zhou, F., Chen, F., and Li, H. Meta-SGD: Learning to learn quickly for few shot learning. arXiv preprint arXiv:1707.09835, 2017.
  36. 36.Liu, P., Yuan, W., Fu, J., Jiang, Z., Hayashi, H., and Neubig, G. Pre-train, prompt, and predict: A systematic survey of prompting methods in natural language processing. arXiv preprint arXiv:2107.13586, 2021.
  37. 37.Nadaraya, E. A. On estimating regression. Theory of Probability & its Applications, 9(1):141–142, 1964.
  38. 38.Olsson, C., Elhage, N., Nanda, N., Joseph, N., DasSarma, N., Henighan, T., Mann, B., Askell, A., Bai, Y., Chen, A., Conerly, T., Drain, D., Ganguli, D., Hatfield-Dodds, Z., Hernandez, D., Johnston, S., Jones, A., Kernion, J., Lovitt, L., Ndousse, K., Amodei, D., Brown, T., Clark, J., Kaplan, J., McCandlish, S., and Olah, C. In-context learning and induction heads. arXiv preprint arXiv:2209.11895, 2022.
  39. 39.Park, E. and Oliva, J. B. Meta-curvature. In Advances in Neural Information Processing Systems, 2019.
  40. 40.Power, A., Burda, Y., Edwards, H., Babuschkin, I., and Misra, V. Grokking: Generalization beyond overfitting on small algorithmic datasets. abs/2201.02177, 2022.
  41. 41.Raghu, A., Raghu, M., Bengio, S., and Vinyals, O. Rapid learning or feature reuse? Towards understanding the effectiveness of MAML. In International Conference on Learning Representations, 2020.
  42. 42.Ramsauer, H., Schäfl, B., Lehner, J., Seidl, P., Widrich, M., Adler, T., Gruber, L., Holzleitner, M., Pavlović, M., Sandve, G. K., Greiff, V., Kreil, D., Kopp, M., Klambauer, G., Brandstetter, J., and Hochreiter, S. Hopfield networks is all you need. arXiv preprint arXiv:2008.02217, 2020.
  43. 43.Rusu, A. A., Rao, D., Sygnowski, J., Vinyals, O., Pascanu, R., Osindero, S., and Hadsell, R. Meta-learning with latent embedding optimization. In International Conference on Learning Representations, 2019.
  44. 44.Schlag, I., Irie, K., and Schmidhuber, J. Linear transformers are secretly fast weight programmers. In ICML, 2021.
  45. 45.Schmidhuber, J. Evolutionary principles in self-referential learning, or on learning how to learn: the meta-meta-... hook. Diploma thesis, Institut für Informatik, Technische Universität München, 1987.
  46. 46.Schmidhuber, J. Learning to control fast-weight memories: An alternative to dynamic recurrent networks. Neural Computation, 4(1):131–139, 1992. doi: 10.1162/neco.1992.4.1.131.
  47. 47.Thrun, S. and Pratt, L. Learning to learn. Springer US, 1998.
  48. 48.Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A. N., Kaiser, L., and Polosukhin, I. Attention is all you need, 2017.
  49. 49.von Oswald, J., Zhao, D., Kobayashi, S., Schug, S., Caccia, M., Zucchet, N., and Sacramento, J. Learning where to learn: Gradient sparsity in meta and continual learning. In Advances in Neural Information Processing Systems, 2021.
  50. 50.Watson, G. S. Smooth regression analysis. Sankhyā: The Indian Journal of Statistics, Series A, pp. 359–372, 1964.
  51. 51.Widrow, B. and Hoff, M. E. Adaptive switching circuits. In 1960 IRE WESCON Convention Record, Part 4, pp. 96–104, New York, 1960. IRE.
  52. 52.Yun, S., Jeong, M., Kim, R., Kang, J., and Kim, H. J. Graph transformer networks. In Wallach, H., Larochelle, H., Beygelzimer, A., d’Alché-Buc, F., Fox, E., and Garnett, R. (eds.), Advances in Neural Information Processing Systems, 2019.
  53. 53.Zhang, A., Lipton, Z. C., Li, M., and Smola, A. J. Dive into deep learning. arXiv preprint arXiv:2106.11342, 2021.
  54. 54.Zhao, D., Kobayashi, S., Sacramento, J., and von Oswald, J. Meta-learning via hypernetworks. In NeurIPS Workshop on Meta-Learning, 2020.
  55. 55.Zhmoginov, A., Sandler, M., and Vladymyrov, M. HyperTransformer: Model generation for supervised and semi-supervised few-shot learning. In Chaudhuri, K., Jegelka, S., Song, L., Szepesvari, C., Niu, G., and Sabato, S. (eds.), Proceedings of the 39th International Conference on Machine Learning, volume 162 of Proceedings of Machine Learning Research, pp. 27075–27098. PMLR, 17–23 Jul 2022. URL https://proceedings.mlr.press/v162/zhmoginov22a.html.
  56. 56.Zucchet, N. and Sacramento, J. Beyond backpropagation: bilevel optimization through implicit differentiation and equilibrium propagation. Neural Computation, 34(12), December 2022.

Citation

MLA
Oswald, J. V., et al. “Transformers Learn In-Context by Gradient Descent”. International Conference on Machine Learning, vol. 202, 2023, pp. 35151–74, https://proceedings.mlr.press/v202/von-oswald23a.html.
APA
Oswald, J. V., Niklasson, E., Randazzo, E., Sacramento, J., Mordvintsev, A., Zhmoginov, A., & Vladymyrov, M. (2023). Transformers Learn In-Context by Gradient Descent. International Conference on Machine Learning, 202, 35151–35174. https://proceedings.mlr.press/v202/von-oswald23a.html
Chicago
Oswald, J. V., E. Niklasson, E. Randazzo, et al. 2023. “Transformers Learn In-Context by Gradient Descent”. International Conference on Machine Learning 202: 35151–74. https://proceedings.mlr.press/v202/von-oswald23a.html.
Harvard
Oswald, J.V. et al. (2023) “Transformers Learn In-Context by Gradient Descent”, International Conference on Machine Learning. PMLR, pp. 35151–35174. Available at: https://proceedings.mlr.press/v202/von-oswald23a.html.
Vancouver
1. Oswald JV, Niklasson E, Randazzo E, Sacramento J, Mordvintsev A, Zhmoginov A, Vladymyrov M (2023) Transformers Learn In-Context by Gradient Descent. In: International Conference on Machine Learning. PMLR, pp 35151–35174

BibTeX

@InProceedings{pmlr-v202-von-oswald23a,
  title = 	 {Transformers Learn In-Context by Gradient Descent},
  author =       {Von Oswald, Johannes and Niklasson, Eyvind and Randazzo, Ettore and Sacramento, Joao and Mordvintsev, Alexander and Zhmoginov, Andrey and Vladymyrov, Max},
  booktitle = 	 {Proceedings of the 40th International Conference on Machine Learning},
  pages = 	 {35151--35174},
  year = 	 {2023},
  editor = 	 {Krause, Andreas and Brunskill, Emma and Cho, Kyunghyun and Engelhardt, Barbara and Sabato, Sivan and Scarlett, Jonathan},
  volume = 	 {202},
  series = 	 {Proceedings of Machine Learning Research},
  month = 	 {23--29 Jul},
  publisher =    {PMLR},
  pdf = 	 {https://proceedings.mlr.press/v202/von-oswald23a/von-oswald23a.pdf},
  url = 	 {https://proceedings.mlr.press/v202/von-oswald23a.html},
  abstract = 	 {At present, the mechanisms of in-context learning in Transformers are not well understood and remain mostly an intuition. In this paper, we suggest that training Transformers on auto-regressive objectives is closely related to gradient-based meta-learning formulations. We start by providing a simple weight construction that shows the equivalence of data transformations induced by 1) a single linear self-attention layer and by 2) gradient-descent (GD) on a regression loss. Motivated by that construction, we show empirically that when training self-attention-only Transformers on simple regression tasks either the models learned by GD and Transformers show great similarity or, remarkably, the weights found by optimization match the construction. Thus we show how trained Transformers become mesa-optimizers i.e. learn models by gradient descent in their forward pass. This allows us, at least in the domain of regression problems, to mechanistically understand the inner workings of in-context learning in optimized Transformers. Building on this insight, we furthermore identify how Transformers surpass the performance of plain gradient descent by learning an iterative curvature correction and learn linear models on deep data representations to solve non-linear regression tasks. Finally, we discuss intriguing parallels to a mechanism identified to be crucial for in-context learning termed induction-head (Olsson et al., 2022) and show how it could be understood as a specific case of in-context learning by gradient descent learning within Transformers.}
}
Metadata:DOI registry

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

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/