Train for the Worst, Plan for the Best: Understanding Token Ordering in Masked Diffusions

18 Jul 2025 · 14 min · 6 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

Compares mass diffusion models (MDMs) vs autoregressive models (ARMs) for generating discrete sequences (text, proteins) and explains why “train for the worst, plan for the best” works via adaptive inference that chooses token fill order using oracle-like strategies.

Guests

No guest names or backgrounds are provided in the transcript; it’s a two-person discussion.

Key claims

MDMs train with order-agnostic masked infilling (harder than next-token prediction) and may underperform on standard likelihood, but adaptive inference can dramatically improve results without retraining by filling positions the model is most certain about.

Notable examples

Sudoku accuracy <7% with random order, ~90% with top-K probability margin; Zebra puzzles ~77% to >98%. Adaptive MDM (6M params) beats larger teacher-forced ARMs (42M–87M). Text perplexity drops with similar diversity; hard Sudoku generalization ~50% vs ARM ~32–33%.

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

Exploring Masked Diffusion Models

0:45 to 1:37

Learn about the main focus of today's discussion: MDMs vs ARMs.

“And you'll see this really interesting tension, MVMs.”

How MDMs Operate

1:37 to 3:56

Understand how masked diffusion models generate content differently.

“That's the core puzzle we're looking at.”

Training Challenges of MDMs

3:56 to 5:13

Examine the complexities and challenges faced during MDM training.

“So the MDM is basically training for literally every possible fill-in-the-blank game you could throw at it, no matter how weird or difficult the pattern of blanks is.”

Adaptive Inference Strategies

5:13 to 10:40

Explore how adaptive inference can enhance MDM performance post-training.

“So MDMs train on these incredibly tough, almost impossible problems sometimes.”

Remarkable Performance Gains

10:40 to 11:16

Witness the dramatic improvements in accuracy using adaptive inference.

“It suggests that this unsupervised way of finding the right reasoning path at inference time can be incredibly powerful, maybe even more powerful than trying to explicitly teach the path during training.”

Generalization and Broader Applications

11:16 to 13:14

Discuss the broader implications of MDMs and their potential applications.

“And what about handling harder problems than it was trained on?”
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:00Okay, so you've probably heard a lot about generative AI lately. You know, models that create text, images, maybe even code. Yeah, it's pretty much everywhere now. These digital brains, they seem to just conjure things up. But have you ever stopped to think how they actually, like, learn to create stuff? It's a really fundamental question. Well, today we're doing a deep dive into exactly that. We're looking at two main ways these models generate things, especially for stuff like text, you know, discrete things, or even proteins. Right, sequences of distinct units, not like smooth, continuous images.

0:34Exactly. And we're going to uncover something pretty surprising about how one approach can seriously outperform the other. We're basically unpacking mass diffusion models, MDMs, and putting them side by side with the more traditional autoregressive models, ARMs. Okay. And you'll see this really interesting tension, MVMs. They train in a super complex way, almost like training for the absolute worst case scenario. Yeah, preparing for anything. But then when they're actually creating something, they have this incredible flexibility. They can sort of plan for the best. That's a great way to put it.

1:07So our mission today is to figure out, is that flexibility at the end during inference enough to make up for that really tough training? And what does that mean for actually solving things like, say, logic puzzles? It's a really interesting tradeoff. Yeah. And this whole deep dive, it's based on a really fascinating paper called Train for the Worst, Plan for the Best, Understanding Token Ordering in Masked Diffusions. A very apt title. Okay. So let's get into it. What if a model that learned the hard way, tackling the toughest problems, what if it could suddenly get way better, way more efficient, just with a clever trick after it's done learning, compared to models that seem simpler?

1:50That's the core puzzle we're looking at. All right. Let's lay out these two approaches first. ARMs versus MDMs. So think about writing a sentence word by word, right? Just left to right. The standard way you'd think about it. That's basically how autoregressive models ARMs work. They predict the next word or token based on all the words that came before it. It's very sequential, like a train on a fixed track, as you said, one direction, step by step. Exactly. A meticulously planned journey, but only on that one track. No detours. And that's where MDMs are just fundamentally different. They don't do the left to right thing at all.

2:28Okay, so how do they work then? Well, MDMs generate stuff by essentially reversing a process of adding noise. Imagine you start with, say, a complete sentence and you randomly blank out or mask some of the words. Okay. The MDM learns to fill in those blanks. But the key is the masking is random. So when it's time to generate, it can start with a totally blank slate and fill in the pieces in pretty much any order it wants, not just left to right. Ah, okay. So it's less like writing a sentence and more like solving a crossword puzzle maybe, filling in bits here and there. That's a pretty good analogy, yeah.

3:01It gives them enormous flexibility when it comes time to actually generate what we call inference time. They can decode tokens, fill in those blanks in basically any order, arbitrary order. That flexibility, it sounds amazing. But, you know, usually great flexibility comes with, well, some kind of catch, right? Sounds like it would be way harder to train. You hit the nail on the head. That's the tradeoff. ARMs, they learn a relatively simple task, predict the next word given the past ones. It's a limited set of problems in a way. Right, just one specific type of prediction. Exactly. But MDMs, they're trained to solve a huge number of different infilling problems, exponentially large, actually.

3:43They have to learn how to predict a missing token based on surrounding tokens that could be anywhere in any configuration. Wow, okay. We call it order agnostic training because the model has to learn to work regardless of the order or position of the known information. So the MDM is basically training for literally every possible fill-in-the-blank game you could throw at it, no matter how weird or difficult the pattern of blanks is. That's the idea. It's preparing for an immense range of scenarios. Which leads to the obvious question. Does trying to learn everything mean it doesn't learn the normal stuff as well?

4:16Yeah. Does this massive training effort actually hurt its performance sometimes? That's exactly what we looked into. And yes, there's evidence, both theoretical and from experiments, that MDMs do train on subproblems that are incredibly hard, even computationally intractable. Intractable, like fundamentally too difficult. For certain configurations, yes. And we saw this play out. For example, with text data, MDMs often underperform compared to ARMs when you measure things like likelihood, basically, how well they predict standard sequential text. Why is that? Well, partly because many of those random masking problems they train on are just much, much harder than the simple next word prediction that ARMs focus on.

4:56It's almost like the MDM is trying to become fluent in every possible variation of a language at once. Well, ARM just focuses on mastering standard grammar and vocabulary. Exactly. So the ARM gets really good at that one thing, while the MDM's effort is spread much thinner across many harder things. Okay. So MDMs train on these incredibly tough, almost impossible problems sometimes. Sounds like a disadvantage, mostly. But you mentioned they have this flexibility at inference time. Is that where the magic happens? Is there a way to use that flexibility to kind of dodge the hard stuff it learned?

5:31This is where it gets really cool. So yes, the training is tough. But the paper shows that with just some simple tweaks at the inference stage, so after training is totally done, no extra learning, you can get these huge performance boosts. Without retraining, just changing how you use the model. Precisely. The key idea is what we're calling adaptive inference. Instead of just randomly picking which blank to fill in next, which is kind of the default way. You get strategic. Exactly. You strategically choose the order. Now think about it. If you had a perfect MDM, one that could solve any infilling problem flawlessly, the order wouldn't matter, right?

6:07Right. It would get the right answer no matter which blank you asked it to fill first. But real-world MDMs aren't perfect. They're better at some infilling problems than others. so the insight is let's guide the model during inference let's have it tackle the problems it's good at first and maybe leave the ones it struggles with until later when it has more context so the knowledge is kind of already baked into the model from that hard training mm-hmm just needs a smarter way to access or apply it you got it that's the core insight it's not about learning more it's about using what it learned more effectively that's fascinating it really changes how you think about what the model knows.

6:46So how do you actually do that? How do you decide which blank, which token position to fill in next? So the paper introduces these things called oracles to guide the process, not like mystical oracles. Uh-huh, right. No crystal balls involved. No, no. Think of them more as intelligent rules or strategies for picking the next step. The basic idea is simple. Try to fill in the positions where the model seems most certain about the answer. Okay. How do you measure certainty? Well, there are a couple of ways. One strategy is called top K probability. You look at all the blank positions and for each one, you see what's the highest probability the model assigns to any possible word or token for that spot.

7:26Then you pick the say K1 position where that highest probability is the biggest. So you fill in the spot where the model has some answer it feels really strongly about, even if you don't know what that answer is yet. Kind of, yeah. But a potentially better way is the top K probability margin. This one's a bit more subtle. Instead of just looking at the single highest probability. You look at the difference. Exactly. You look at the difference between the probability of the most likely token and the second most likely token for that position. Ah, okay. So if the model is like, it's definitely the with 90 % probability and the next best is A with 2%, the margin is huge.

8:04High certainty. Right. But if it's saying, hmm, it could be cat with 45 % probability or maybe dog with 40 % probability, the margin is tiny. Indicating it's really unsure, confused between those two. Precisely. So the margin strategy says, pick the position with the biggest gap between the first and second choice. Avoid those confusing spots until you have more information. Go for the easy wins first. And crucially, these strategies are deciding which position to fill next, not what word to put there, right? No. The model still decides the actual word based on its probabilities. It's a critical distinction, yes.

8:39We're just guiding the order of decisions, not the decisions themselves. Okay, this sounds clever. But does it actually work? What happens when you apply this adaptive inference to real tasks? The results were, well, pretty stunning, actually. Especially on tasks that require some logic or constraint satisfaction. Like the logic puzzles you mentioned, Sudoku. Exactly, take Sudoku. Using a standard, non-adaptive MDM, just letting it fill things in random, the accuracy was terrible. Like, less than 7%. 7%. That's barely better than guessing. Maybe worse. Pretty much unusable. But then, using the same trained model, just switching the inference strategy to that top K probability margin, accuracy jumped to nearly 90%.

9:20Wait, hold on. From 7 % to 90%. Just by changing the order, it fills in the grid. Just by changing the order. No retraining, same model. For another type, Zebra Puzzles, which are also logic-based, it went from about 77 % to over 98%. That is absolutely wild. A 10x or more improvement in accuracy just from being smarter about the inference path. It really highlights how much potential was locked up inside the model, just waiting for the right way to be accessed. And how does this compare to the ARMS, the left-to-right models? That's maybe even more fascinating. So this adaptive MDM didn't just beat the standard ARMs.

9:57It even outperformed specialized ARMs. Specialized how? ARMs that were specifically trained using something called teacher forcing to learn the correct sequence for solving the puzzle. So they were explicitly guided during training on the right path. Okay, so they had a big advantage in training. A huge advantage. And these ARMs often had way more parameters like the building blocks of the model. We saw cases where a 6 million parameter MDM using the adaptive strategy beat ARMs with 42 million or even 87 million parameters that were trained with teacher force. Wow. So a smaller model trained in that general order agnostic way could outperform much larger specially trained models just by using a smarter inference strategy.

10:40Yes. It suggests that this unsupervised way of finding the right reasoning path at inference time can be incredibly powerful, maybe even more powerful than trying to explicitly teach the path during training. That really turned some common assumptions on their head. What about text? Did it help there, too? It did. We saw significant drops in perplexity. That's a standard measure for language models. Lower is better means it's less surprised by the text it sees or generates. So it got better at predicting text. Yes, noticeably better. And importantly, it did this while keeping roughly the same level of diversity in the generated text.

11:13It didn't just collapse into predicting repetitive stuff. Okay, that's crucial. And what about handling harder problems than it was trained on? Generalization. Another really interesting finding, we trained models on, say, easy and medium Sudoku puzzles and then tested them on hard ones they'd never seen. How'd they cope? Well, accuracy dropped for all models, which you'd expect, but the adaptive MDMs were significantly more robust. They still managed almost 50 % accuracy on the hard puzzles, while the best ARM dropped to around 32-33%. So the MDM's difficult, broad training actually helped it generalize better when faced with unexpected difficulty.

11:53It seems that way. That train for the worst approach, grappling with all those diverse and hard infilling problems, seems to help it extract a deeper, more fundamental understanding of the problem structure, which pays off when things get tougher. So let's wrap this up. This deep dive, it really shows something quite profound, doesn't it? These mass diffusion models, MDMs, they go through this incredibly tough training, training for the worst, tackling this huge space of complex problems. Yeah, a really demanding learning process. But that same process gives them this hidden superpower. flexibility at inference.

12:26And if you unlock that with a smart strategy like these adaptive oracles plan for the best approach. You get these dramatic performance leaps. Yeah. Suddenly they can solve logic puzzles with amazing accuracy, outperform models that are way bigger or had special training, and even generalize better. All just by guiding the order in which they apply their knowledge. It's a powerful demonstration of separating the training process from the inference strategy. It really makes you think. If models can get so much better just by applying what they already know more strategically at inference time, where else could this apply?

13:01What other really complex problems could benefit from this kind of plan for the best adaptive approach, especially if the underlying learning process is messy or seems inefficient that trains for the worst style? You know, think about challenges in logistics, maybe planning complex deliveries or schedules, or scientific discovery, sorting through possibilities, even creative fields. Yeah, anywhere where the path to the solution isn't obvious or fixed, where finding the right order of steps is key, it could be very broadly applicable. What stands out to you, the listener, when you think about that?

13:34Where else could strategically navigating the possibilities unlock hidden potential? Something to ponder.

From the publisher

This academic paper explores masked diffusion models (MDMs), a promising approach for generative modeling in discrete domains. It investigates the trade-off between training complexity and inference flexibility in MDMs compared to autoregressive models (ARMs). The authors demonstrate that MDMs are trained on computationally challenging subproblems, leading to performance imbalances. However, they show that adaptive inference strategies, which strategically select the token decoding order, can significantly enhance MDM capabilities, allowing them to circumvent these difficult problems. Notably, adaptive MDMs achieve superior performance on logic puzzles like Sudoku, even surpassing ARMs with more parameters and explicit training for decoding order.

More from Best AI papers explained

All 475 episodes
Train for the Worst, Plan for the Best: Understanding Token Ordering in Masked DiffusionsBest AI papers explained · 14 min
Listen in VO