Learning to reason in LLMs by expectation maximization

28 Dec 2025 · 14 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

The episode explains a research paper that frames LLM “chain-of-thought” training as a latent-variable model approximated by expectation maximization (EM). It introduces filtered EM (FEM): sample a rationale, reward it with a binary check (correct answer = 1 else 0), and update only from successful rationales. Key claim: the main driver of improved reasoning is the rationale proposal distribution, not complex training pipelines.

Notable examples

a broken-bone question where PPS-trained rationales become concise and correct (cell division) versus messy pretraining rationales.

Guests

none mentioned; the episode is a host-led Deep Dive.

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 Chain of Thought Prompting

0:41 to 2:42

Explains the importance of rationale in large language models and introduces the expectation maximization algorithm.

“Our mission is to unpack how the core mechanism of LLM fine-tuning can be seen as an approximation of a very old algorithm, the expectation maximization algorithm, or EM.”

Latent Variable Models and Logical Analysis

2:42 to 4:53

Discusses the mathematical formalization of reasoning and the introduction of latent variables in LLMs.

“The EM algorithm, it's designed for exactly these situations where you have unobserved variables like zivil dollars.”

Expectation Maximization Algorithm Breakdown

4:53 to 7:20

Details the E-step and M-step of the EM algorithm and their application in LLMs.

“It relies on a super simple binary reward function.”

Challenges of Applying EM to LLMs

7:20 to 8:57

Explains the computational challenges of directly applying the EM algorithm to large language models.

“But if it fails, star moves to its key rationalization stage.”

Filtered EM: An Approximate Solution

8:57 to 11:12

Introduces the filtered EM (FEM) approach as a practical solution for training LLMs.

“Doesn't that create a model that's a great rationalizer, but a poor problem solver?”

Comparing Rationale Generation Strategies

11:12 to 13:16

Compares different strategies for generating rationales, including rejection sampling and prompt posterior sampling.

“produce a more compact, more focused justification.”

Final Thoughts and Implications

13:16 to 13:52

Discusses the implications of using hints in training and the nature of learning in models.

“applied in the right way, generates the best results.”
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 back to the Deep Dive. We all know that large language models are, well, just immensely powerful. They can generate everything from poetry to complex code. But when you ask them to tackle a truly hard problem, say a difficult multi-step calculation or nuanced exam question, their performance often hinges on one crucial thing, showing their work. We call this intermediate process the rationale or chain of thought prompting. If the model can generate a coherent step-by-step justification, it usually finds the right answer. But just asking for a rationale isn't enough. The real challenge is how do we actually teach an LLM to generate high-quality reasoning consistently?

0:36That is the multi-billion dollar question, and it's what we're diving into today. This research takes this very modern LLM challenge, how to train better reasoning, and grounds it really firmly in classical statistics. Our mission is to unpack how the core mechanism of LLM fine-tuning can be seen as an approximation of a very old algorithm, the expectation maximization algorithm, or EM. And by doing that, the researchers could rigorously compare different ways of generating these internal rationales, these thoughts, and figure out which method is the most efficient path to a smarter model. And our source material for this is a paper titled Learning to Reason in LLMs by Expectation Maximization.

1:15It provides this wonderful theoretical bridge connecting these common reward-based fine-tuning methods we see everywhere with the formal structure of EM. Okay, let's unpack this. We're going to start with the formal language. How do the researchers actually view reasoning mathematically? To get this beyond just showing your work, they formalize it through the lens of a latent variable model, an LVM. Think of it like this. If you give the model a difficult question, let's call it$6, and ask for a final answer,$1, dollars, that direct jump is just too hard. The logical leap is too great. Exactly.

1:47It's too dense. So this difficulty means you need to introduce an unobserved intermediate step. We call that the rationale. This dollar is the latent variable. It's the chain of thought, the logical analysis, all the internal work. So you're breaking the impossible direct task,$6, into two much easier stages. Precisely. The structure becomes a Markov chain. The question six-dollar generates the rationale, and then the rationale leads to the correct answer dollar. And crucially, the model, conditioned on both the original question and those carefully generated steps, has a much, much higher chance of getting that correct answer than if it just tried to guess.

2:23That makes immediate sense. The rationale isn't just window dressing. It's the critical ingredient that captures the logic. We're not just training the model to guess the answer. we're training it to build a set of steps that justifies the answer. We're building verifiable logic paths. Which brings us straight to the tool used in classical statistics to solve these kinds of problems. Expectation maximization. The EM algorithm, it's designed for exactly these situations where you have unobserved variables like zivil dollars. And it works by alternating between two famous steps. First, you have the E step, the expectation step.

2:58This computes the posterior distribution over that latent variable ZIL, given the data you can see, so solar cells dollars. In our world, that means calculating what the optimal rationale should look like, considering all possible paths that lead to the right answer. Okay. Then you have the M-step, the maximization step. This updates the model's parameters to maximize the probability of generating those ideal rationales you just found. That's very elegant. The E-step figures out the perfect thought process, and the M-step teaches the model to actually have that thought process. But I'm guessing that as soon as you apply this to a modern LLM, that E-STEP becomes a computational nightmare.

3:34You are absolutely right. The exact EM algorithm is completely infeasible for LLMs. And here's why. That E-STEP expectation requires calculating the probability distribution over all possible rationales. That could lead from the question to the answer. And for a, say, 3 billion parameter model, the number of possible rationales, the number of possible sequences of text is effectively infinite. It is. It's like trying to find the average of an infinite set of numbers. You just can't do it. You can't sum over that infinite space. So that's why so many LLM training methods feel like clever hacks.

4:10They're trying to get around this possible calculation. Exactly. And this is the PISIT point in the research where they translate the ideal theory into something you can actually run on a GPU. So what's the compromise? How do they make EM work? They adopted an approximate solution they call filtered EM, or FEM. And FEM simplifies both the E-step and the M-step pretty dramatically. First, the E-step. Instead of calculating that infinite expectation, it's approximated by a single Monte Carlo sample. The model just generates one rationale answer pair, a cat. So it stops trying to calculate everything and just tries one path.

4:44Yes. That immediately cuts the computational burden from infinite down to one. Okay, that handles the E step. What about the M step? The M step then uses what they call a filtered gradient update. This is the critical filtering part. It relies on a super simple binary reward function. If the sampled answer matches the real answer job, the reward is one. If not, it's zero. So let me translate that for our listeners. The filtered EM process is basically saying, generate a step-by-step thought process, check if it leads to the right answer. If it does, great, we'll learn from that thought process.

5:16If it lands on the wrong answer, we just throw that thinking away and learn nothing from it. Is that the core logic? That is the core logic. The model is only fine-tuned on the rationales. The Haases successfully lead to the correct answer. And the theoretical significance here is massive. The paper shows that this iterative filtered process, which is what a lot of successful self-improvement algorithms already do, can be viewed directly as a mathematically grounded approximation of classical EM. It gives a theoretical backbone to methods that were previously just treated as empirical tricks that happen to work.

5:50Right. Okay, so if the filtered EM structure is fixed, you sample, check, filter, and fine-tune, then the performance difference between all these different self-improvement methods has to come down to one thing, how you generate that initial rationale sample. That's it. In this FEM framework, the key engineering decision is the rationale proposal distribution. They call it to dollar. This dollar is the strategy the model uses to generate those initial rationale candidates. And you need a good dollar, one that balances two competing goals. Okay, what are they? First, high success probability. You need the model to generate correct answers often, so you have plenty of training data.

6:25Second, high rationale quality. The logic needs to be good, so the model generalizes well later when you test it without any help. Right, okay. So let's compare the three main schemes they tested for generating this dollars. Let's start with the baseline, rejection sampling with budget M, RSM. RSM is the most straightforward kind of brute force approach. The model just samples rationale answer pairs conditioned only on the question,$6. It just keeps sampling up to a maximum budget, say M on$5. The moment it stumbles upon the correct answer, it saves that rationale for fine-tuning. If it fails five times, it gives up on that question for that round.

7:03Just keeps rolling the dice. Okay. Then you have star, self-taught reasoner, which was some pretty seminal work in this area. It's designed to be smarter than just random guessing. Star is a more complex two-stage policy. It starts with a minimal attempt, a single rejection sample, so RS with a budget of one. Correct. If that single, unprompted attempt gets the right answer, that rationale is used. But if it fails, star moves to its key rationalization stage. In this second stage, it generates a rationale using a prompt that explicitly reveals the correct answer as a hint. This, you know, significantly increases the odds of producing a usable rationale, especially early on when the model isn't very good at reasoning on its own.

7:43So STAR is a fallback. It's like, try it on your own, but if you fail, I'll give you the answer and then you show me the proof. And this brings us to the new approach that turned out to be the surprising winner in this research. Prompt, posterior sampling, or PPS? PPS is powerful because it's so simple and efficient. The researchers realized that the most effective part of STAR was that second stage, the part with the hint. So PPS just gets rid of the first stage entirely. It discards the complex rejection sampling. So it never even tries on its own during training. Never. PPS only uses that rationalization prompt, which conditions the generation on the true answer, high dollars, as a hint.

8:19It does it right from the very first sample every single time. Which means PPS immediately tells the model the goal. It bypasses all the trial and error and says, look, the answer is D. Now generate the most coherent steps that prove D is correct. That's it. And the way they do it is really just in the prompt itself. The prompt will literally include text like hashtag hashtag hint. The best answer is answer D. Your reasoning should explain and justify it. It just leverages the model's ability to follow instructions. Wait, I have to stop you there because that sounds a lot like cheating during training.

8:52You're just giving the answer away. Doesn't that risk the model learning to just justify whatever answer you give it instead of learning to reason for itself? Doesn't that create a model that's a great rationalizer, but a poor problem solver? That is the core tradeoff they acknowledge, and it's a fantastic question. Yes, you are introducing a mismatch between training where it gets the hint and testing where the hint is gone. but the data suggests the massive increase in the yield of correct high-quality rationales, the teaching material, more than makes up for that mismatch. By using the hint, PPS generates teaching material that is highly structured and guaranteed to be on the right track.

9:31It's an efficient way to bootstrap reasoning by using what the model's already good at, which is following instructions. It's prioritizing the quality of the internal instruction set over the difficulty of creating it. Okay, let's look at the testing setup. So they ran tests on 3 billion parameter instruction-tuned models, LAMA 3.2 and QUEN 2.5, across three big multi-choice benchmarks, ARC, MMLU, and OpenBook QA. And they applied this filtered EM process for five iterative updates, which is enough to see a pretty significant change in performance. And the results really validate the PPS approach.

10:06The figures in the paper are pretty clear. PPS consistently got the highest final test accuracy across all benchmarks and on both models. After five iterations, it wasn't just beating simple projection sampling. It was also outperforming the more complex, chained approach of STAR. And this leads to the efficiency win, which I think is the most insightful part. You might assume PPS won because, you know, by giving the answer away, it generated more total training examples. That would be the logical assumption. Easier process, more data, better results. But when they measured the data usage, So the actual number of correct rationales used for fine-tuning STAR often used more total training data than PPS.

10:45Wait, really? How? Because STAR combines that first rejection sampling attempt with the hint stage, so it could get a correct answer from either stage. So STAR used more successful training examples overall, yet PPS achieved higher test accuracy with less data. That's the critical distinction. It implies that the rationales generated by PPS aren't just more numerous, they're inherently higher quality. more effective for fine-tuning. By conditioning on the correct answer, the model is directed to produce a more compact, more focused justification. It doesn't waste time exploring all these logical dead ends that rejection sampling might.

11:20And you can really see that qualitative difference in the examples they provide. There's one where the question is, a broken bone heals through the process of. And the options include things like cell division, adaptation, etc. Right. And if you look at the rationale the model generates before fine-tuning it's a mess it just rambles for three steps it considers adaptation it considers chemical digestion and then it lands on the wrong answer it's totally unfocused but then you compare that to the reasoning it produces after being trained with filtered EM using PPS the change is dramatic the final rationale is concise it's accurate it skips the rambling it immediately frames the question around tissue repair and correctly identifies cell division as the fundamental process.

12:04The logic is just sharp and directed. It turned the model from a confused rambler into a focused expert. So PPS drove the model toward a much more functional and reliable internal representation of the solution. Exactly. Okay, let's pull all these threads together. For us, for the listeners who are either building these things or just curious about how they learn, what is the ultimate insight here? I think the core insight is that by formalizing LLM reasoning, with this classical statistics lens of EM, the research confirmed that the design of the rationale sampling distribution, the teaching strategy, is the most important part of self-improvement.

12:39It's not about designing these super complex sequential policers or even just maximizing the raw number of correct examples. It's about efficiently generating the highest quality, most accurate rationales you can during training. And the surprising practical takeaway is that the simplest, most direct approach Prompt posterior sampling wins. It bypasses all that trial and error by just giving a hint, and in doing so, it leverages the model's instruction-following abilities to manufacture perfect teaching material. It takes a power the model already has and turns it into a self-teaching mechanism to build better internal reasoning.

13:15It shows that sometimes the simplest intervention, applied in the right way, generates the best results. That leaves us with one final provocative thought to consider that builds right on top of this. If explicitly revealing the correct answer during training leads to superior internal reasoning and better test performance later on, what does that actually tell us about the nature of learning in these models? Is the optimal path to true, independent reasoning always through explicit conditioning on the goal? Forcing the model to reverse engineer the perfect justification rather than expecting it to discover the path all on its own through trial and error, something for you to mull over until our next deep dive.

From the publisher

This research formalizes the process of reasoning in large language models as a latent variable model, utilizing the expectation-maximization (EM) algorithm to improve performance. The authors demonstrate that training a model to generate intermediate rationales before answering is mathematically equivalent to reward-weighted fine-tuning using binary correctness as a signal. A central focus of the study is the sampling distribution used to create these rationales, comparing methods like rejection sampling and the self-taught reasoner (STaR). The paper introduces prompt posterior sampling (PPS), a technique that conditions the model on the correct answer during training to generate more effective reasoning traces. Experiments across multiple benchmarks show that PPS consistently outperforms existing methods by producing more concise and accurate rationales. Ultimately, the work highlights that high-quality rationale generation is just as critical to model improvement as the underlying optimization algorithms.

More from Best AI papers explained

All 475 episodes
Learning to reason in LLMs by expectation maximizationBest AI papers explained · 14 min
Listen in VO