Transformers Learn to Implement Multi-step Gradient Descent with Chain of Thought

7 Mar 2026 · 17 min · 7 chapters

Ask about this episode

Ask anything about it. ChatGPT or Claude reads this page and answers with the times it was said.

Connect VO and ask about every podcast you hear, including the moments you saved. Add to ChatGPT · Add to Claude

In short

How chain-of-thought prompting changes a transformer’s optimization behavior, letting it implement multi-step gradient descent rather than a single-step update; includes related “loop transformer” experiments.

Guests

None named in the transcript (two hosts discuss the paper).

Guest/paper backgrounds (from episode)

Jianhao Huang, Zixuan Wang, Jason Dealey (authors of an ICLR 2025 spotlight paper).

Key claims

A one-layer linear transformer trained for direct prediction can only realize one gradient-descent step from zero weights, causing high error. Training with chain-of-thought (intermediate predictions) makes the model unroll T gradient steps autoregressively, using generated tokens as a “memory register” for updated weights. The learned algorithm generalizes out-of-distribution (unseen covariance matrices, noise variance, and weight magnitudes). Loop transformers (tied weights iterated across depth) achieve similar multi-step unrolling without visible tokens.

Notable examples

In-context linear regression via attention computing a covariance-like inner-product structure; multi-step recovery of ground-truth weights; out-of-distribution covariance/noise/weight scaling tests; comparison to loop transformers.

Written by AI. May contain mistakes. Listen to the episode to check what was said.

Chapters

Tap a time to open that second in VO

Understanding Large Language Models

0:45 to 2:14

Exploring how language models operate and their reasoning capabilities.

“And it really is the necessary next step in our understanding of these architectures.”

Unpacking the ICLR 2025 Paper

2:14 to 4:43

Examination of a paper revealing the mechanics behind chain of thought prompting.

“So instead, the authors set up a very focused experiment.”

Single Layer Limitations and Breakthroughs

4:43 to 6:41

Discussion on the limitations of a single-layer model and how chain of thought overcomes these.

“So, the researchers introduced the chain of thought objective to the exact same architectural setup.”

Comparing Iterative Approaches in Transformers

6:41 to 11:08

Comparative analysis of chain of thought and loop transformers in optimization.

“through the autoregressive generation of intermediate steps.”

Scaling Concepts to Complex Models

11:08 to 12:19

Examining the implications of findings for scaling to larger, complex models.

“The loops transformer achieves the multi-step gradient descent entirely internally, without the overhead of generating visible tokens.”

The Impact of Prompting Techniques

12:19 to 14:00

Understanding how prompting rewires optimization pathways in models.

“So the specific math gets more complicated.”

Understanding Multi-Step Gradient Descent

14:00 to 16:42

Learn how chain of thought prompts enhance transformer model performance.

“And understanding that difference is what allows researchers to build more efficient architectures.”
Hear the part that matters, and keep it.Open this episode in VO. Double tap your headphones to save a moment as you listen.
Get VO free

Transcript

Automatic transcript. May contain errors.

0:00Welcome to today's Deep Dive. I want you to just take a second and think about the underlying mechanics of the large language models you interact with daily. Right. Yeah. The ones we all use. Exactly. I mean, we all know the standard operational procedure by now, right? Right. You feed the model a complex logic problem. It struggles. So you append that magic phrase, think step by step. And suddenly the output is flawless. Right. Suddenly it works perfectly. And you know that it works. We all do. But today, we are taking you under the hood of these architectures. We want to understand how they actually learn and reason mathematically when they're prompted this way.

0:39It's really moving beyond just the empirical observations. Yeah, we're looking straight into the hidden mechanics of their optimization pathways. And it really is the necessary next step in our understanding of these architectures. I mean, we've spent the last few years just marveling at the emergent capabilities of autoregressive generation. Absolutely. But it is an entirely different level of comprehension to look at the actual mathematical proofs governing the gears turning inside the black box. And to help us see those gears, we are unpacking a highly revealing ICLR 2025 spotlight paper by researchers Jianhao Huang, Zixuan Wang, and Jason Dealey.

1:17It's a fantastic paper. It really is. The mission for our deep dive today is to explore exactly how giving an AI a chain of thought physically changes its underlying mathematical behavior, how it moves it from a constrained state to an incredibly expressive one. What's fascinating here is that despite the remarkable empirical success of chain of thought prompting and its obvious theoretical advantages in making models more expressive by granting them additional compute tokens, well, the actual mechanisms underlying chain of thought training have remained largely unexplored. Right, from a rigorous theoretical standpoint anyway.

1:51Exactly. It has been a persistent mystery. We know the what of the performance boost, but the how has been hiding in the optimization dynamics until this research. Okay, let's unpack this. Because to figure out the how, the researchers needed a transparent testing ground. I mean, analyzing a trillion-parameter mixture of experts model while it's writing poetry isn't going to give us clean mathematical proofs. No, definitely not. There's way too much noise. Right. So instead, the authors set up a very focused experiment. They used an in-context weight prediction task for linear regression. Which is an incredibly elegant choice for isolating the variables.

2:29By stripping away the massive nonlinearities and focusing on in-context linear regression, they can directly observe the learning dynamics of the attention mechanism. Because they're essentially asking the model to act as an optimization algorithm on the fly. Exactly. Given a sequence of X and Y data points in the prompt, plus a new query X, the model has to calculate the underlying weight vector to predict the final Y. And because you already know how linear regression and gradient descent operate, Right? The fascinating part here is how the researchers isolate the transformer's capacity to execute these algorithms.

3:04Right. They specifically analyze a one-layer linear transformer to serve as the baseline. And they establish a critical architectural limitation right out of the gate, don't they? They do. When you give this baseline model the in-context learning problem and just ask for the final answer. The standard direct prediction objective. Yeah, exactly. When you do that, it can only implement a single step of gradient descent. Wow, just one step. Just one. And that structural constraint is exactly why the direct prediction baseline fails on complex tasks. Mathematically, the self-attention operation in a single linear layer computes the inner products of the queries and keys.

3:42Which effectively builds a sample covariance matrix from the prompt's context. Right. It then applies the value matrix to output a prediction. And this entire operation perfectly mirrors a single gradient update step, starting from a zero-initialized weight vector. Which means it mathematically guarantees a high error rate on anything but the most trivial data, because you can't optimize a high-dimensional space in a single leap. You really can't. It takes one step down the gradient, hits its architectural limit, and stops. It has no mechanism to iteratively refine that weight vector. It highlights a hard algorithmic bottleneck because reaching the optimal weight vector, the true minimum of the loss landscape, requires multiple iterative updates.

4:27A one layer model simply doesn't have the depth to unroll those iterations internally. It computes its single update, projects the prediction and completely fails to recover the ground truth weight vector. Here's where it gets really interesting. We have this mathematically proven failure state for the single layer baseline. So, the researchers introduced the chain of thought objective to the exact same architectural setup. This is the turning point. Right. Instead of demanding the final prediction immediately, they train the model to output a sequence of intermediate predictions before the final answer.

5:02And this is the absolute breakthrough of the study. When the transformer is trained with this chain of thought objective, it bypasses its structural limitations entirely. It learns to perform multi-step gradient descent. Yes. And it executes this auto-regressively. Let's drill down into that autoregressive execution, because this is where the genius of the prompt structure meets the math. By forcing the model to generate an intermediate output token, you are effectively giving it a memory register. A summary register. That phrasing captures the dynamic perfectly. Because the token isn't just text, right?

5:35It's a physical instantiation of the updated weight vector at step T. Exactly. The physical act of generating the intermediate token allows the model to store its current progress. During the next forward pass, the attention mechanism attends to that newly generated token alongside the original context. It uses the stored state to compute the gradient again, applies the learning rate, and outputs the next token. And that new token represents the weight vector at step t plus one. It is literally using its own context window as a computational scratch pad to run an unrolled optimization loop. It's incredible to see the math prove it out.

6:11It's bypassing the depth constraint of its one-layer architecture by unrolling the computation across time. Or across tokens, to be precise. And the contrast in performance is absolute. By utilizing this autoregressive, multi-step gradient descent, the exact same one-layer transformer achieves near-exact recovery of the ground truth weight vector. So the single-step failure transforms into a multi-step success solely because the optimization pathway was rerouted through the autoregressive generation of intermediate steps. That's it entirely. The researchers actually provide the theoretical bounds for this, showing how the training loss drops significantly as the number of chain of thought steps increases.

6:53It proves that the model isn't just getting slightly better. It is fundamentally executing a different order of algorithmic complexity. Right. It moves from an O1 optimization process to an OT process, where T is the number of intermediate tokens. Which brings us to a critical threshold in the study, because demonstrating this on the training distribution is one thing, but proving that the model had internalized the algorithm requires pushing it out of distribution. Yes, generalization is the real test. Right, because the lingering skepticism with LLMs is always whether they are just highly sophisticated pattern matchers interpolating within their training data.

7:29You know, does this just work on the data the AI already memorized? We need to know if the model memorized the trajectory or if it learned the generalized rule of gradient descent. And this raises an important question, and the paper tackles it head on. They presented the chain-of-thought trained model with entirely unseen data structures. We're talking about out-of-distribution covariance matrices, different variance levels in the noise. And scaled magnitudes of the underlying weight vectors that were never encountered during training. And if it were just pattern matching, the autoregressive loop would quickly derail.

8:02Yeah, the errors would compound with every token generation, setting the prediction completely off a cliff. But the exact opposite occurs. The trained transformer effectively generalizes across these out-of-distribution parameters. That's wild. It continues to successfully perform the multi-step gradient descent and consistently converges on the near-exact weights. That is massive. It proves the transformer has learned a repeatable generalized algorithm. It literally learned the mechanics of gradient descent itself, implemented through the autoregressive generation of tokens. It is a profound shift from empirical observation to verifiable algorithmic execution.

8:40The model learns to configure its query, key, and value matrices such that the self-attention operation consistently mirrors a robust gradient update step. Regardless of the specific data values fed into the context window. Exactly. Now, I want to pivot slightly to another architectural concept the researchers explored, which beautifully complements these chain of thought findings. They introduce loop transformers into the experimental setup. Ah, yes. The loop transformer experiments are a fantastic structural counterpoint to the autoregressive findings. And I think bridging these two concepts provides a really complete picture of what is happening mathematically.

9:17Let's bring this up. A standard deep transformer passes the hidden state through sequential distinct layers. Right. but a looped transformer takes the output of a single layer and feeds it back into the exact same layer. With the exact same tied weights. Multiple times. It iterates in depth rather than in sequence length. And the underlying philosophy is identical to what we just discussed with Chain of Thought. We established that standard attention in a single layer equates to one gradient step. If you want more steps, you need more compute. Chain of Thought buys that compute by unrolling the process across generated tokens.

9:52While looped transformers by that compute by unrolling the process iteratively through the same set of weights. And the researchers demonstrated that when you apply this looping mechanism to the in-context learning of linear regression, you see a massive spike in final performance compared to transformers without looping. Because the mathematical result is essentially the same unrolled gradient descent. They map perfectly onto each other. So whether you're providing the network with the space to execute iterative steps via the context window sequence length. Or by the iterative depth of a looped layer.

10:26You are solving the same expressivity bottleneck. You're giving the optimization algorithm the computational cycles it needs to converge on the optimal weights. Which really makes you look at the massive parameter counts of today's frontier models differently. Oh, totally. If a model can iteratively refine its internal representations using shared weights or intermediate tokens, adding billions of parameters across distinct layers might actually be computationally inefficient compared to simply allowing a smaller model to loop or think longer. That is the core takeaway for architecture design moving forward.

11:02The paper proves that computational expressivity isn't strictly bound by parameter count. It is bound by the available pathways for iterative refinement. The loops transformer achieves the multi-step gradient descent entirely internally, without the overhead of generating visible tokens. But it relies on the exact same fundamental principle of iterative state updates. I do want to play devil's advocate for a moment, though. We are talking about a one-layer linear transformer solving linear regression. It's a highly idealized convex optimization problem. How confident are we that these exact mechanics scale to the highly nonlinear billion-parameter behemoths optimizing incredibly complex non-convex lost landscapes in natural language tasks?

11:47That's a vital distinction to make. We cannot assume a direct one-to-one mapping of these specific linear proofs to a nonlinear network. like GPT-4 or LAMA. Right. However, foundational proofs like this serve as the theoretical bedrock. While the exact geometry of the gradient steps in a deep nonlinear network will be vastly more complex... The fundamental mechanism absolutely holds true. Exactly. Using intermediate output tokens as computational scratch pads to bypass the architectural depth constraints of a single forward pass, that holds true. So the specific math gets more complicated. But the architectural paradigm, the transition from a single step guess to a multi-step iterative refinement, is the universal principle at play.

12:28Precisely the point. This study provides the mathematically rigorous proof of concept that has been missing. It anchors our empirical observations and verifiable optimization theory. If we pull back for a second, it feels like this research fundamentally shifts how we should visualize the prompting process. Oh, completely. We often talk about prompt engineering as if we were trying to speak the model's language, you know. like coaxing it into giving us a better answer through linguistic tricks. If we connect this to the bigger picture, the implications of this paper completely dismantle that linguistic perspective.

13:02Prompting isn't a conversational hack. No. When you introduce a chain of thought objective, you are actively rewiring the optimization pathway. You are physically altering the model's capabilities by granting it the requisite computational space to execute iterative mathematical updates. You are providing the structural scaffolding for the math to occur. And without that scaffolding, without the intermediate tokens acting as memory registers, the model is physically constrained from reaching the correct output. No matter how perfectly you phrase the final question. Right, because the model is architecturally bound to its single-step forward-pass limit.

13:39The fact that we can mathematically prove this transition, that we can actually map the multi-step gradient descent happening as the intermediate tokens are generated, It moves our understanding of large language models away from empirical alchemy and solidifies it as rigorous mathematical chemistry. It gives you a profound appreciation for the sheer mechanics of it all. Yes. It's not magic. It's unrolled optimization. And understanding that difference is what allows researchers to build more efficient architectures. It allows practitioners to leverage these tools with actual intent rather than just trial and error.

14:14So let's recap the journey we've just been on together. We started with the ubiquitous knowledge that appending a chain of thought prompt dramatically improves a model's performance on logic and reasoning tasks. But today, thanks to the deep dive into this paper, we map the mechanical reality underneath that performance boost. We established that a standard transformer model, constrained to a direct prediction objective, gets stuck. It attempts to solve an optimization problem with a single step of gradient descent, effectively hitting an architectural wall and failing to recover the correct weights.

14:46But the introduction of the chain of thought objective fundamentally alters that dynamic. It empowers the transformer to utilize the intermediate tokens as state representations. Allowing it to perform multi-step gradient descent auto-regressively. Exactly. It feeds its own computational progress forward, transforming that initial single step failure into a robust, generalized algorithm that flawlessly handles out-of-distribution data. Furthermore, we saw that internal architectural changes, like the shared weights of loop transformers, achieve a mathematically similar unrolling of the optimization process.

15:22Proving that iterative processing, whether across time or across depth, is the key to complex problem solving. Absolutely. So what does this all mean for you? Whether you're building applications on top of these models or simply tracking the trajectory of AI development, it underscores a crucial reality. The structure of how we ask the model to process information dictates its mathematical ceiling. Your prompt is the algorithm's operating environment. Yes. By enforcing intermediate steps, you aren't just asking for clarity. You are literally provisioning the compute cycles required for high-dimensional optimization.

15:58I want to leave you with a final thought to mull over, building directly on the mechanics we've just unpacked. We have seen mathematically that forcing a transformer to output its intermediate states allows it to jump from a single step failure to a multi-step success in predicting exact weights. The tokens are the bridge across the computational bottleneck. Exactly. So if merely requiring the generation of intermediate tokens unlocks this hidden mathematical capacity for complex optimization, what other fundamental limitations of current AI architectures might be solved, not by adding trillions of new parameters, but simply by re-architecting the constraints on how the machine is allowed to iteratively refine its internal state before delivering a final output?

16:41man that is a phenomenal question to end on it completely reframes the race for bigger models versus the push for smarter unrolled computation thank you so much for joining us on this deep dive we hope you walk away not just knowing that your ai works but armed with a precise understanding of the mathematical elegance happening right beneath the surface we will catch you on the next deep dive

From the publisher

This research paper explores how Chain of Thought (CoT) prompting enables transformers to solve complex mathematical problems by mimicking iterative optimization techniques. The authors demonstrate that while standard models are limited to a single stage of calculation, using intermediate reasoning steps allows a transformer to execute multi-step gradient descent internally. Through the lens of linear regression tasks, the study proves that this autoregressive process leads to a near-perfect recovery of underlying data patterns that simpler models cannot capture. Furthermore, the findings indicate that looped architectures and CoT significantly boost the ability of these models to generalize to new information. Ultimately, the work provides a formal theoretical framework to explain why breaking down problems into smaller parts enhances the algorithmic power of large language models.

More from Best AI papers explained

All 475 episodes
Transformers Learn to Implement Multi-step Gradient Descent with Chain of ThoughtBest AI papers explained · 17 min
Listen in VO