Bypassing the Exponential Dependency: Looped Transformers Efficiently Learn In-context by Multi-step Gradient Descent

Bo Chen, Xiaoyu Li, Yingyu Liang, Zhenmei Shi, Zhao Song

INTRODUCTION

Large Language Models (LLMs) have gained immense success and have been widely used in our daily lives, e.g., GPT4 , Claude 3.5 , and so on, based on its Transformer architecture . One core emergent ability of LLMs is In-Context Learning (ICL) . During ICL, the user provides an input sequence (prompts) containing some question-answer pairs as in-context examples and the goal query that the user cares about, where the examples and query are drawn from an unknown task. The LLMs can in-context learn from these examples and generate the correct answer for the goal query without any parameter update. Notably, these unknown tasks may never seen by LLMs during their pre-training and post-training, e.g., synthetic task . Thus, many believe that the in-context learning mechanism is different from supervised learning or unsupervised learning, where the latter may focus on feature learning, while the ICL may perform some algorithm to learn. For instance, the Transformer can implement algorithm selection or gradient descent by in-context learning forward pass.

Many works have tried to understand how Transformers perform single-step gradient descent . They study one-layer linear Transformer in-context learning on linear regression tasks and show that the one-layer linear Transformer can perform single gradient updates based on the in-context examples during a forward pass. Recent work has further shown that a linear looped Transformer can implement multi-step gradient descent updates with multiple forward passes, meaning that Transformers can express multi-step algorithms during in-context learning. However, their theoretical results require an exponential number of in-context examples, exp⁡(Ω(T))\exp(\Omega(T)), where TT is the number of loops or passes, to achieve a reasonably low error for linear regression tasks. This violates the intuition that more gradient descent updates lead to better performance.

Thus, it is natural to ask the following question:

Is it necessary to use an exponential number of examples for Transformers to implement multi-step gradient descent during in-context learning?

In this work, we study linear looped Transformers (Definition 3.6) in-context learning on linear vector generation tasks (Definition 3.4), which is as hard as linear regression. We show that the linear looped Transformer can efficiently perform multi-step gradient descent as long as in-context examples are well-conditioned. We present our main result in the following theorem.

In Theorem 1.1, as long as the condition number is constant, we can see that the linear looped Transformer will perform better when the loop number is increasing, i.e., the error will exponentially decay to . Informally, for a small constant ϵ\epsilon, if we draw XX from Gaussian distribution, the condition number of X⊤XX^{\top}X will be 1≤κ≤1+O(ϵ)1\leq\kappa\leq 1+O(\epsilon) when n≥Ω(d/ϵ2)n\geq\Omega(d/\epsilon^{2}). Thus, informally, we only need O(d)O(d) numbers of in-context examples to guarantee a good performance. Furthermore, our preliminary experiments (Section 6) validate the above arguments and our theoretical analysis.

The main intuition of our analysis is that we find that a linear looped Transformer can explicitly perform gradient descent in its hidden states (Lemma 4.1 and Theorem 4.2). Thus, the error analysis can be directly solved by the standard convex optimization technique (Theorem 5.8).

We study linear looped Transformers in-context learning on linear vector generation tasks (Definition 3.3), which is as hard as linear regression.

We find that linear looped Transformers can explicitly perform gradient descent in their hidden states (Lemma 4.1 and Theorem 4.2).

We demonstrate that linear looped Transformers can efficiently perform multi-step gradient descent as long as in-context examples are well-conditioned, e.g., n=O(d)n=O(d), where the error will exponentially decay to after TT loops (Theorem 4.4).

Our preliminary experiments on synthetic data validate our main theoretical results (Section 6).

Roadmap. This paper is structured as follows: we begin with a review of related work in Section 2, followed by essential definitions and foundational concepts in Section 3. Section 4 delves into the gradient computation analysis within the Looped Transformer architecture, examining both individual layer computations and the full looped structure. In Section 5, we analyze the error convergence of looped transformers under strong convexity and smoothness conditions, providing an upper bound on error after TT gradient descent iterations for linear vector generation. Section 6 presents our experimental results and findings. Finally, we conclude in Section 7.

RELATED WORK

This section briefly reviews the related research work on Large Language Models (LLM), In-Context Learning (ICL), and looped transformers. These topics have a close connection to our work.

Neural networks based on the Transformer architecture have swiftly become the dominant paradigm in machine learning for natural language processing applications. Expansive Transformer models, trained on diverse and extensive datasets and comprising billions of parameters, are called large language models (LLM) or foundation models . Examples include BERT , PaLM , Llama , ChatGPT , GPT4 , among others. These LLMs have demonstrated remarkable general intelligence capabilities across various downstream tasks.

Researchers have developed numerous adaptation techniques to optimize LLM performance for specific applications. These include methods such as adapters , calibration approaches , multitask fine-tuning strategies , prompt tuning techniques , scratchpad approaches , instruction tuning methodologies , symbol tuning , black-box tuning , reinforcement learning from the human feedback (RLHF) , chain-of-thought reasoning and various other strategies. And here are more related works aiming at enhancing model efficiency without compromising performance, such as .

In-context Learning.

A significant capability that has emerged from Large Language Models (LLMs) is in-context learning (ICL) . This feature allows LLMs to generate predictions for new scenarios when provided with a concise set of input-output examples (referred to as a prompt) for a specific task, without requiring any modification to the model’s parameters. ICL has found widespread application across various domains, including reasoning , self-correction , machine translation and many so on.

Numerous studies have focused on enhancing the ICL and zero-shot capabilities of LLMs . A substantial body of research has been dedicated to investigating the underlying mechanisms of transformer learning and in-context learning through both empirical and theoretical approaches. Building upon these insights, our analysis extends further to elucidate ICL implementing multi-step gradient descent.

Looped Transformer.

The concept of recursive inductive bias was first introduced by into Transformers. Looping Transformers are also related to parameter-efficient weight-tying Transformers theoretically showed that the recursive structure of Looped Transformers enables them to function as Turing machines. demonstrate that increasing the number of loop iterations can improve performance on some tasks. Recent studies have provided theoretical insights into the emulation capabilities of specific algorithms and their convergence during training, with a particular emphasis on in-context learning. However, the expressive capacity of Looped Transformers or Looped Neural Newrok in function approximation and their associated approximation rates remain largely unexplored territories.

PRELIMINARY

In this section, we present some preliminary concepts and definitions of our paper. In Section 3.1, we introduce some basic notations used in our paper. In Section 3.2, we defined some important variables to set up our problem.

2 In-context Learning

First, we introduce some definitions of in-context examples and their labels.

Note that the θ∗\theta^{*} is unseen to the model. Then, our vector generation tasks are defined as follows.

We remark that our vector generation task is as hard as the linear regression task, as they are dual problems. Solving any one of them requires to estimate the θ∗\theta^{*}.

Combining all above, we have the ICL prompt/input data for the model.

3 Linear Looped Transformer

In line with recent work by GSR+ and ACDS , we consider a linear self-attention model, formally defined as follows:

where M=[0n×n0n×111×n0]M=\begin{bmatrix}\mathbf{0}_{n\times n}&\mathbf{0}_{n\times 1}\\ \mathbf{1}_{1\times n}&0\end{bmatrix} is a casual attention mask for text generation.

In Definition 3.5, we combine the query matrix and key matrix as QQ and combine the value matrix and output matrix as PP for simplicity following previous works .

Building upon this, we introduce the concept of a Linear Looped Transformer :

Let TT be the loop number. Let η(t)>0\eta^{(t)}>0 for any t∈{0,1,…,T−1}t\in\{0,1,\dots,T-1\}. The linear looped transformer TF(Z(0);Q,P)\mathsf{TF}(Z^{(0)};Q,P) is defined as

The Looped Transformer is to simulate a real multiple-layer Transformer with residual connections , where η\eta represents the weights of the residual components.

Our settings are more practical than in the following sense.

We have a more practical casual attention mask used in generation. Our mask requires Hardamard product, the same as the standard attention, while uses matrix product mask.

Our model does not have prior knowledge of XX distribution, while the model in knows the distribution of XX, i.e., distribution free. In practice, the LLMs do not know any information about in-context examples.

4 Linear Regression with Gradient Descent

In this section, we introduce some key concept of linear regression with gradient descent.

In this work, we consider n>dn>d, assuming X⊤XX^{\top}X is inevitable. Then, we have θ∗=θ~\theta^{*}=\widetilde{\theta}.

In the realizable setting (no noise term), if n>dn>d, then θ∗=θ~\theta^{*}=\widetilde{\theta}.

Finally, we define the condition number, which will be used in our final convergence bound.

We define the condition number of input data as κ:=λmax⁡(X⊤X)λmin⁡(X⊤X)\kappa:=\frac{\lambda_{\max}(X^{\top}X)}{\lambda_{\min}(X^{\top}X)}.

GRADIENT COMPUTATION IN LOOPED TRANSFORMER

In this section, we present a comprehensive analysis of the gradient computation process within the Looped Transformer architecture. Our investigation begins with an examination of computations in individual layers and subsequently extends to the full looped structure. This approach allows us to build a nuanced understanding of the Looped Transformer’s behavior, starting from its fundamental components.

We commence our analysis by establishing a crucial result regarding the output of a single layer in our Looped Transformer model. This foundational lemma serves as a cornerstone for our subsequent derivations and provides valuable insights into the model’s inner workings.

Let Z(0)Z^{(0)} be defined in Definition 3.4. Let Q=Id+1,d+1Q=I_{d+1,d+1}. Let P=[Id×d0d×101×d0]P=\begin{bmatrix}I_{d\times d}&\mathbf{0}_{d\times 1}\\ \mathbf{0}_{1\times d}&0\end{bmatrix}. Let causal attention mask be M=[0n×n0n×111×n0]M=\begin{bmatrix}\mathbf{0}_{n\times n}&\mathbf{0}_{n\times 1}\\ \mathbf{1}_{1\times n}&0\end{bmatrix}. Then, we have

where the first step follows from Definition 3.5, the second step follows from Definition 3.4, and the rest steps directly follow from the matrix multiplication. ∎

This result illuminates the specific form of the attention mechanism’s output in a single layer, which is essential for understanding the model’s overall behavior, where the output only has non-zero terms in the position of q(0)q^{(0)} and the format is close to Eq. (1).

Having characterized the behavior of a single layer, we now extend our analysis to encompass the full-looped transformer structure. Turning our attention to the core of our analysis with the groundwork laid for single-layer computations. The following theorem establishes a crucial relationship between the transformer’s output and the iteratively updated parameters:

We will prove this theorem by induction on tt, where t∈[T]t\in[T]. From Lemma 4.1, we have:

Where the first step follows from Definition 3.6, the second step follows from Definition 3.4 and Definition 3.5, the third step follows from basic algebra. Then we extract q(0)⊤−η(0)(q(0)⊤X⊤X+αy⊤X)q^{(0)\top}-\eta^{(0)}(q^{(0)\top}X^{\top}X+\alpha y^{\top}X), we have

where the first step follows from basic algebra, the second step follows from we defined θ(0)=−1αq(0)\theta^{(0)}=-\frac{1}{\alpha}q^{(0)}, the third step follows from basic algebra. Thus, q(1)=−αθ(1)q^{(1)}=-\alpha\theta^{(1)}, so θ(1)=−1αq(1)\theta^{(1)}=-\frac{1}{\alpha}q^{(1)}.

Similarly, by math induction, we can have θ(T)=−1αq(T)\theta^{(T)}=-\frac{1}{\alpha}q^{(T)}. Thus, we finish the proof by TF(Z(0);Q,P):=−q(T)\mathsf{TF}(Z^{(0)};Q,P):=-q^{(T)}. ∎

To further refine our understanding of the Looped Transformer’s performance, we introduce a bound on the final prediction error:

where the first step is from Definition 3.3 and Theorem 4.2, the second step follows ∥θ∗∥2=1\|\theta^{*}\|_{2}=1, the third step is from the linear properties of inner product, and the fourth step is from Cauchy-Schwarz inequality. ∎

This result provides a quantitative measure of the model’s accuracy, linking it directly to the number of iterations and multi-step gradient descent results of linear regression in Eq. (2).

Finally, we present our main theoretical contribution, which encapsulates the core findings of our work:

Let κ\kappa be the condition number defined in Definition 3.11. Let TT be the number of loops. Let the initial point q(0)=0dq^{(0)}={\bf 0}_{d}. Then, we have the final prediction error satisfies

The proof directly follows Lemma 4.3 and Theorem 5.8. ∎

In Theorem 4.4, as long as the condition number is constant, we can see that the linear looped Transformer will perform better when the loop number is increasing, i.e., the error will exponentially decay to . Usually, we only need O(d)O(d) numbers of in-context examples to guarantee a constant κ\kappa.

The above theorem offers a comprehensive characterization of the Looped Transformer’s behavior, providing a tight bound on the prediction error that decays exponentially with the number of iterations. The intuition is that the Linear Looped Transformer can explicitly perform gradient descent in its hidden states. Furthermore, our theoretical finding is also supported by our experiments in Section 6.

Under condition 8Td2n≤122T\frac{8Td^{2}}{\sqrt{n}}\leq\frac{1}{2^{2T}}, we have the optimal linear regression error is ≤8Td222Tn\leq\frac{8Td^{2}2^{2T}}{\sqrt{n}}.

In their work, the linear looped transformer has an error bound 8Td222T/n{8Td^{2}2^{2T}}/{\sqrt{n}}, while our bound is ∣α∣⋅exp⁡(−T2κ)|\alpha|\cdot\exp(-\frac{T}{2\kappa}). As the looped number TT increases, our error bounds will exponentially decay, while theirs is exponentially increase. Note that our linear vector generation task and their linear regression task are dual problems. Our results align with the common intuition that more steps of gradient descent lead to better performance.

ERROR CONVERGENCE

In this section, we explore the convergence properties of looped transformers, focusing on their behavior under conditions of strong convexity and smoothness. We begin by defining these key concepts and then proceed to establish their implications.

We first introduce some crucial definitions.

We say that μ\mu is the strong convexity constant of ff.

To quantitatively analyze the parameter dynamics in our linear vector generation task, we first derive the Lipschitz and convexity constants for the model introduced in Definition 3.8.

where LL is the Lipschitz constant defined in Definition 5.2.

where the first step follows from Definition 3.8, and the second step follows from properties of norm. From Definition 5.2, we observe that

where ∥X⊤X∥\|X^{\top}X\| is the spectral norm of X⊤XX^{\top}X denoting the maximum eigenvalue. ∎

The following two lemmas are closely related and build upon each other to establish the strong convexity constant for a specific optimization problem.

where μ\mu is the strong convexity constant defined in Definition 5.1

where the first step follows from Definition 3.8, the rest step follow from basic algebra and Lemma 5.4. From Definition 5.1, we observe that

where λmin⁡(X⊤X)\lambda_{\min}(X^{\top}X) denotes the minimum eigenvalue of X⊤XX^{\top}X. ∎

2 Main Result

We first commence with a statement of Lemma 5.6, which furnishes a convergence rate for gradient descent on strongly convex and smooth functions.

We now present a rigorous upper bound on the error magnitude of the gradient descent algorithm’s output after TT iterations, elucidating the convergence properties of this optimization method in the context of linear vector generation.

Let the condition number κ=λmax⁡(X⊤X)λmin⁡(X⊤X)\kappa=\frac{\lambda_{\max}(X^{\top}X)}{\lambda_{\min}(X^{\top}X)}.

Let μ\mu and LL be defined in Definition 5.1 and Definition 5.2.

The initial point θ(0)\theta^{(0)} satisfies ∥θ(0)−θ∗∥2≤R\|\theta^{(0)}-\theta^{*}\|_{2}\leq R.

where the first step follows from Lemma 5.6, the second step follows from we choose η(t)=1L\eta^{(t)}=\frac{1}{L}. Then consider the term μL\frac{\mu}{L}, we have

where the first step follows from Lemma 5.3 and Lemma 5.5, the second step follows the definition of condition number κ=λmax⁡(X⊤X)λmin⁡(X⊤X)\kappa=\frac{\lambda_{\max}(X^{\top}X)}{\lambda_{\min}(X^{\top}X)}.

where the first step follows from basic algebra, the second step follows from (1−1/κ)κ<e−1(1-1/\kappa)^{\kappa}<e^{-1} (Fact 5.7), and the last step follows from ∥θ(0)−θ∗∥2≤R\|\theta^{(0)}-\theta^{*}\|_{2}\leq R. ∎

Theorem 5.8 tells us that GD can well solve linear regression tasks. In particular, when the input data has a good condition number, the approximation error will exponentially decay to . We use the above insights in the proof of our main results (Theorem 4.4).

EXPERIMENTS

In this section, we aim to verify our theory by evaluating the convergence behavior of gradient descent for linear vector generation. We designed our experiment to examine the impact of varying sample sizes on convergence rates while keeping the feature dimension fixed. Our results demonstrate that empirical convergence rates consistently outperform theoretical upper bounds across all sample sizes, with significant improvement in convergence speed as the condition number decreases, validating our theoretical predictions.

Result Interpretation.

Our experiment investigates the convergence behavior of gradient descent for linear vector generation with varying sample sizes nn and a fixed feature dimension d=4d=4. Figure 1 illustrates the convergence rates for different nn values {16,32,64,128}\{16,32,64,128\}, comparing empirical results with theoretical bounds in Theorem 4.4. A key aspect of this experiment is the condition number κ\kappa, which decreases as the number of examples increases. The average κ\kappa values for the different sample sizes are {4.57,3.13,2.07,1.62}\{4.57,3.13,2.07,1.62\}, corresponding to n={16,32,64,128}n=\{16,32,64,128\} respectively. This inverse relationship between nn and κ\kappa is noteworthy, as it significantly influences the convergence rates. The results demonstrate that as the sample size nn increases, the convergence rate improves substantially. This is evident from the steeper slopes of both empirical and theoretical lines for larger nn values. Importantly, the empirical convergence rates consistently outperform the theoretical upper bounds across all sample sizes, with the gap between empirical and theoretical performance narrowing as nn increases. This observation aligns with our theoretical expectations in Theorem 4.4 and highlights the crucial role of the condition number in determining convergence behavior.

CONCLUSION

In this work, we have demonstrated that linear looped Transformers can efficiently implement multi-step gradient descent for in-context learning, requiring only a reasonable number of examples when input data is well-conditioned. This finding relieves the previous assumptions of an exponential number of in-context examples and offers new insights into the capabilities of Transformer architectures. Our theoretical analysis and preliminary experiments pave the way for more efficient inference algorithms in large language models and open avenues for future research in this domain.

Acknowledgments

Research is partially supported by the National Science Foundation (NSF) Grants 2023239-DMS, CCF-2046710, and Air Force Grant FA9550-18-1-0166.

References