In short
Scaling Up Test-Time Compute with Latent Reasoning with Jonas Geiping - #723
Podcast Overview
- Podcast Title: The TWIML AI Podcast
- Host: Sam Charrington
- Guest: Jonas Geiping, Research Group Leader at Ellis Institute and Max Planck Institute for Intelligent Systems
- Episode Focus: Discussion on Geiping's recent paper, “Scaling Up Test-Time Compute with Latent Reasoning: A Recurrent Depth Approach”
- Key Concepts: Reasoning models, recurrent depth architecture, internal reasoning vs. verbalized reasoning, dynamic compute allocation, and model evaluation.
---
Key Discussions
Background on Jonas Geiping
- Experience: Background in mathematics and computer science.
- Research Focus: Safety and efficiency aligned learning, particularly in recurrent models to solve complex problems.
Paper Overview
- Main Proposal: Introduces a novel language model architecture that uses recurrent depth to enable "thinking in latent space."
- Comparison of Reasoning: Explores internal reasoning (not verbalized) versus verbalized reasoning, drawing parallels to human cognitive processes.
Recurrent Depth Architecture
- Concept: Unlike traditional recurrent neural networks (RNNs) which process sequences, this architecture allows for a flexible number of computational layers during inference, enhancing reasoning capabilities.
- Mechanism:
- Uses a small number of layers that can be repeated.
- Each iteration refines the model's hidden state, allowing it to think "deeper" without discarding previous computations.
Motivation for the Approach
- Context of Research: Emerged during a period of increased interest in reasoning models and efficiency in training language models.
- Comparison to MoE Models: This architecture uses fewer parameters but recycles them through recurrent processing, contrasting with Mixture-of-Experts (MoE) that utilize more parameters less often.
Training Process
- Initial Model Size: Started with a 100 million parameter model, later scaled to a 3.5 billion parameter model.
- Challenges: Faced complexities in distributing training across a supercomputer, necessitating custom solutions for efficient operation.
Findings and Results
- Performance: The model exhibited strong reasoning capabilities, outperforming other models—specifically excelling in tasks like grade school math and coding challenges.
- Dynamic Compute Utilization: At test time, the model adapts its compute resources based on the complexity of tokens, displaying an emergent behavior that allows it to specialize in tasks during inference.
---
Key Takeaways
Internal vs. Verbalized Reasoning
- Human Analogy: Just as humans can think through problems without verbalizing, this model allows for deeper reasoning without generating intermediate text outputs.
Architectural Advantages
- Efficiency: The recurrent depth approach simplifies large language models, making them more memory-efficient and potentially faster during inference.
- Adaptive Exits: The model can learn when to stop processing based on the convergence of hidden states, allowing for dynamic adjustment of compute resources.
Implications for Future Research
- Model Safety: Considerations around transparency and interpretability in AI systems, especially as models begin to "think" internally rather than always outputting text.
- Further Experiments: Future work focuses on fine-tuning and exploring how to optimize data mixes for improved performance.
---
Closing Remarks
- The episode provides valuable insights into innovative approaches to language model architecture and training, emphasizing the importance of efficient reasoning mechanisms and their implications for AI development. Geiping’s work highlights the potential of recurrent depth models to advance the current landscape of AI reasoning capabilities.
For complete show notes and additional resources, visit [TWIML AI Podcast Episode 723](https://twimlai.com/go/723).
Written by AI. May contain mistakes. Listen to the episode to check what was said.
Transcript
Automatic transcript. May contain errors.0:00I think it's clear that we can sometimes think longer about a problem without actually verbalizing the answer in our head. We don't always think in steps in our head. But even as humans, we have these two axes in which we can scale, compute or thinking in this way. And I think that's kind of interesting, right? It's not necessarily you can do this or the other. But maybe as humans, we naturally do both of those. And these models also could do both.
0:45All right, everyone, welcome to another episode of the TwiML AI podcast. I'm your host, Sam Charrington. Today, I'm joined by Jonas Skyping. Jonas is a research group leader at Ellis Institute and the Max Planck Institute for Intelligent Systems in Tübingen. Before we get going, be sure to take a moment to hit that subscribe button wherever you're listening to today's show. Jonas, welcome back to the podcast. Yeah, happy to be back. A lot has happened in the last year. A lot has happened in the past year. I'm really excited about the conversation we're about to have. You recently published a paper called Scaling Up Test Time Compute with Latent Reasoning, a Recurrent Depth Approach that fits in at a unique time when, you know, reasoning models are getting popular, you know, just after the quote-unquote deep seek moment, if we want to call it that.
1:36And this paper proposes a different approach to reasoning. So really looking forward to digging into that work. I'd love to have you spend a few minutes talking a little bit about your background and refreshing listeners who may not have heard you introduced last year. Yeah, I'm Jonas. I have an old background in mathematics, but I've been doing computer science maybe since 2016. and I've been a bit around in Germany and the US. Now I'm back in Germany and I'm leading a little research group in Germany in Tübingen where we work on safety and efficiency aligned learning. And these are really things that interest me a lot.
2:19And this project in particular is one that's really been a favorite of mine and really have been, this has been a long time that we've worked on this because we really wanted to figure out how can we train these models very differently? Or we had this idea that we could train these models very differently and it would still be interesting and still be a good model. And so because this whole backstory to this whole recurrent approach is that way back when some good colleagues of mine trained these models in 2020, 2021, we trained these recurrent models to solve mazes. I'm not sure if you have seen that.
3:02It's really a toy problem, right? But it's kind of a very intriguing algorithmic task because you train these small recurrent models on mazes of sizes like 13 by 13. So it's really, it's a tiny pixelized maze, right? And it's like, you look at it, it's so trivial what the solution is. So you train these models on these 13 by 13 mazes. And then at test time, you just run them for more recurrences. so you do you put in more compute and in a test time you can solve arbitrarily large mazes like up the way like you have they have these funny pictures in the paper where it's like uh you know the whole screen is filled like every pixel is a different pixel in this maze it's like a i don't know 8000 by 8000 pixel maze and the model just puts in a lot of compute and just solves it right and it's an interesting example because there this model really has learned the algorithm to solve mazes or to solve this particular kind of maze.
3:57Right. And I think that was very motivating for us. So you see that like you could really learn an algorithm like that. But it was always a bit restricted to like, you know, you could only learn either you learn only mazes or you learn only addition or only chess problems. These were all a bit like maybe toy tasks or very like very limited domains where could see that these approaches would scale. And what I should also say is that, so what's interesting here about recurrence here or how that comes in is I think people are very familiar, or maybe this is maybe the ML person speaking, maybe are very familiar with recurrent neural networks.
4:41Right? And this is a bit different. In a recurrent neural network, what you have is that you have a sequence and the recurrence is basically that like every time you go forward in the sequence, you do another step of computation. And this is an efficient way to compute things about sequences. But what we're doing here is actually a bit different. What we're here doing is we're doing a recurrence in really in the depth of the model. And a good analogy here is that the model doesn't have a fixed amount of layers, but it has a small number of layers that it repeats. And this relates to lots of approaches in that field, like our weight sharing, sometimes called our weight tying.
5:21in this way the model really has like an arbitrary number of layers which really means that in comparison to a recurrent neural network where you always produce more input or more output to do more compute same way in a transformer every time you want to do more compute you have to produce more outputs with this sort of like level of this recurrent depth you can produce an arbitrary amount of compute before you compute the next output so it's sort of like the the compute axis is disentangled from the output axis or from the context of the model and that's what's so interesting i mean algorithmically about this kind of approach got it got it so the immediate thought in hearing the title um and the idea of recurrence is to think oh maybe you're doing something with rnns or like splicing that kind of recurrence into a transformer it's not really like that it's um manipulating the number of layers at test time to kind of flex the amount of compute that's being uh applied to a given generation yeah exactly exactly it's a bit um it relates a bit to these analogies for maybe for an rnn or for transformer where you um
6:44For example, I think the analogy for an RNN would be that you actually don't move to the next token in the sequence, but you sort of like produce like filler outputs and then you keep on iterating on these filler outputs. That would be like where you really stop on the sequence and then you, on your point where you currently stopped, you keep on like recurring a bit and then you move forward. You mentioned how you had this paper that explored this recurrent approach as kind of background. But then I'm wondering, talk a little bit about kind of the impetus to apply that. It comes at an interesting time when, you know, we've got these reasoning models like the O1 series from OpenAI.
7:28You know, DeepSeek showed how you could take a very different approach to training those. You know, this paper came out shortly after that, but I'm suspecting you've been working on it for quite a while. So kind of situate it for us in time and, you know, what motivated you to go after this particular approach of reasoning? Yeah. So basically, this is all the way back in 2023, like Tom and me, Tom is the last author, we had just been working on a paper on training language models more efficiently. And then we've been talking a lot about these recurrent models and how they would be cool to do in language.
8:08And because we really think it would be an interesting language model that could also solve things like math and code better, if it could reason more. Because back then, people were also doing a lot of MOE models. And MOE is interesting because they're the other extreme. It's a model that has more parameters, but uses less of them. So that seems like optimal to maybe store information. And this is, in some conceptual way, it's the opposite. It's a model that has fewer parameters. So fewer parameters, but uses them over and over and over. Yes, exactly. Right. So we really wanted to figure out how we could make this paradigm of recurrent depth work with language and how we could get, how we can modify it, what a good like training objective would be to really scale this up.
8:54Because it really wasn't really shown before in Transformers at all. Like, or there's some, at least like from our, like from this like maze direction, these were always convolutional models. And it wasn't so clear how to put this in Transformer and be nice. There's some interesting parallel work on looped transformers, which goes in a related direction. But it didn't quite work as well for us, just scaling those up larger. And that relates a bit to this part that this actually also is, so like one part that's different here is that it's latent, where basically we actually have a few non-recurrent layers in the model.
9:34And then we have this recurrent part in the middle. And then we have a few more non-recurrent layers at the end. To really have this computation be in this centralized latent space. And then still have some layers that actually decode back into language. And have some layers at the start that decode from language. And then with the right objective, this really helped us a lot. And then we had some, let's say like 100 million parameter prototypes in 2023. and we're like okay what next we're academics right so we we teamed up with lots of people that know a lot more about um high performance computing and we wrote this big allocation grant for um insight is like some big allocation grant basically this compute that we then got at oak rich at oak rich national labs and but then like it's still this um it was quite an interesting journey because it's it's one thing to like train this model in pytorch and it's 100 million parameters and it's kind of nice right to actually seeing okay now we now there are 4 000 accelerators and it still has to work and they have to work together and this like funny little probabilistic objective that you wrote still has to scale and it's just to work here and it was a pretty interesting journey that i personally enjoyed a lot although it was a bit harrowing at times as well to actually make this work and run this on these cards right because we also um so this is um this was compute on frontier and frontier is the largest amd supercomputer in the world and just and then we were just given that amount of compute because the um the folks that inside were saying yeah here we have this you could have this and so um but uh 2024 was also the year where a lot of things were moving at amd and were just beginning to be supported for example uh flash attention like flash attention or something it's like a very basic algorithm for us nowadays but it took like a lot of time maybe like until like half the year maybe until like last july where i felt this was really was stable enough that we could reliably deploy it and run it and there was a lot back and forth and amd did a lot and we had some, maybe like a lot of issues, but we got around them and it did train in the end.
11:54But it was a whole journey maybe of like, maybe a year where the underlying little model didn't really change, right? Or like the, I think in the paper, that's like, yeah, right. Like in the paper, maybe that's like section three and section three is like, okay, 2023. And then all of 2024 is just section four, which is, oh, we trained it on 4 ,000 GPUs. Oh, wow. And that was kind of an interesting journey. And yeah, then it also got pretty exciting because, as you mentioned, we had these test-sam-computed reasoning models coming out in fall, and we still hadn't finished training the model. And we were saying, oh, oh, oh.
12:35And suddenly, this was in the news everywhere, and suddenly this was happening, and we still hadn't finished training the model. And then now we did finish training the model in early December last year, and that's also why we pretty much wrote up our results in January and published them. But it was a bit of a scary moment as the world caught up to us a bit in that, although that's also why the approach is a bit different. Yeah, to what degree did the way you position the results around reasoning, you know, was that impacted by the timing of when you finished the work? like was that specifically the way you were thinking about it going in or was it more like you've got this architecture let's try to scale it and see what it can do and then lo and behold like it produced results that were relevant to you know the zeitgeist at the time it was always this idea of reasoning and learning algorithms but maybe like that that framing is maybe like the more classical framing would be that it's learning an algorithm just be like this older line of papers would maybe call this learning an algorithm or a prior towards learning algorithmic reasoning to call this all just blanket reasoning maybe that's a more modern way of looking at it and the other part it's also very modern is that like in the intro now we clearly call this out that there's a um i think we made this term up but i think it fits quite well to compare sort of like verbalized reasoning which we which is how we how we um maybe the box So we put in these other approaches to what we're doing.
14:09And that's a modern distinction that we put in to claim, okay, what are we actually doing differently? And I really hadn't thought about it so much that it wasn't so clear to me that just long chain of thoughts could scale so well. But maybe this wasn't clear to many of us a year ago. But I think it's interesting to look at them in this dichotomy that one is really trying to learn these algorithms in a high dimensional space. so like before it really decodes into language and the other approach is really trying to scaffold reasoning on to language let's maybe maybe even back up a step and really dig into that because i think that that is an idea that um you know people are really excited about and that is really exciting this idea that uh with traditional language models you know the, you know, the high compute and reasoning that we're seeing, it's all like having the models think out loud.
15:09And in a sense, what this paper is proposing is that we can also have the models kind of think to themselves and then, you know, not necessarily generate, you know, all of that thinking as language. And so with that in mind, kind of talk about how you think about, you know the relationship between those two and you know the the relative merits of one versus the other and and how they might compare like i'm curious like how you think about all that yeah so maybe the first thing i think is interestingly i think this is quite natural for us as well maybe really as humans this is a bit of anthropomorphizing i hope that's fine but basically um i think it's clear that we can sometimes think longer about a problem without actually verbalizing the answer in our head.
15:58Like we don't always like think in steps in our head, especially if it's a problem that maybe relates to maybe spatial reasoning or maybe something that you have some intuitive, maybe motor planning or like planning reasoning or how you go from one place to another. These are often things that we think about and we also think about longer, but we don't really verbalize this kind of thinking. And then interestingly, we also like these models, we can still write things down, right? We can sort of extend our thinking process by verbalizing it, by writing things down in notebooks, in equations. In this way, we can externalize thinking and think even longer.
16:35But even as humans, we have these two axes in which we can scale, compute or thinking in this way. And I think that's kind of interesting, right? It's not necessarily you can do this or the other. But maybe as humans, we naturally do both of those. and these models also could do both. But what's interesting about thinking without verbalizing it is that maybe like from a very basic perspective, this just seems much more powerful because basically the model, now if we come a bit to like the theory of like what algorithms can transformers represent, basically the transformer does some calculation because it is like n layers deep.
17:25Let's say that's like 96 layers deep. So the transformers has some input. The model does some calculation and the next word comes out. And then in the next layer, it can maybe access these previous calculations, but again, only up to the 96th layer. And if the model were to do a calculation that's deeper than 96, it's kind of hard because it has to go into these previous attention parts. so like the the model is a bit limited in its depth you're kind of calling out this idea that like i mean in some ways it's parallel to the argument for test time compute in general it's like that argument is often stated as well you know look at us as humans like there are things that we you know think about and the answer comes to us very quickly and there are other things that we spend a lot of time thinking about.
18:19And so why should models have like just a fixed, you know, compute, fixed budget for, you know, response, which kind of led to the idea of test time compute to some degree and allowing models to spend more time on that compute. And what you're kind of suggesting is that there's like another dimension of that. And there's, there's another the dimension of it from a transformer architecture perspective like you're both time and compute limited and also architecturally limited and you the recurrence kind of gives you some flex in that other dimension i don't know like what the implications of that are but um it strikes me that there are some parallels there yeah yeah i think an interesting analogy to make this clearer is that in some way if we like extract a bit away from it what the transformer is doing if it's doing a very long chain of thought is that it has some hidden representation that are computed on what the next token should be.
19:25Then that whole hidden representation, which is a high dimensional vector, is thrown away and then one token is sampled out of it. And the next step, that token is embedded again to compute a whole other hidden representation, which is thrown away. And then a new token is embedded. Because we keep embedding and unembedding these representations that the model is providing. And in some way, maybe the only thing that this recurrent model is doing is that it's just not doing that. It's basically keeping, it keeps on recurring and keeps on refining this hidden representation instead of always like encoding and decoding it in every step.
20:05And so how does it, how does it know how much to do this, when to stop doing it, that kind of thing? Yeah. So interestingly, that's actually not something we have, we have done a lot in this work. So what we've done a lot is basically is show how to scale these models and how to train them to have this capability. But we actually found that the best way we could scale these models was just to train them with a randomized depth or randomized number of occurrences. Really? Okay. So basically there's some older work also called Universal Transformers, right? Because this whole thing relates to universal Turing completeness.
20:45and these models had halting sub-modules that were doing training you had to decide when to stop but those are pretty hard to scale those are pretty hard to train because you're kind of training a model to get better and at the same time you're training it to exit early and it's hard to bootstrap a model that has to get better and harder problems and has to know when to exit early and so one way we got around this was basically that like during training we have a log normal Poisson distribution. It's really a heavy tail distribution over possible steps. And we sample one of them and we just compute loss on this step and then we back propagate.
21:26And this is how we train the model. But then interestingly, at inference time, what we found is that we could just zero shot adaptive exits. For example, just by computing how much the state changes from step to step. We could just make a simple rule based on that and just exit based on the simple rule that we just put in zero shot afterwards. It was kind of exciting that this actually did work at all. It wasn't a guarantee that this. What's the specific state that you're applying that algorithm to? So we've tried this with several ones. We've tried this with the actual latent recurrence state and for example the L2 norm of that state to the previous state.
22:08But we've also done this in the output distribution. so the right we have like a well maybe we're at state 16 and so we compute the next token distribution and we can then we go to state step let's say 24 we can compute the next token distribution and then we can compare the colbyg-like divergence of these two probability distributions and if that is small then we just exit the idea being that the model is converging on some thought or idea exactly and interestingly like the second one seems a bit more practically better because what we observe is sometimes the model keeps on thinking in this in this recurrent space and we have some interesting graphics in the paper where like what is actually like how we can project that down in lower dimensions how it's thinking there but this thinking doesn't always reflect in the probability of the next token right that's sort of like a more stable state is actually what do you predict next right and it's like it's not it wasn't so clear to us beforehand, or maybe also just wasn't clear in general, whether the model would learn this specialization.
23:12Basically, this was a research question, was like, whether just with scale, would the model learn to specialize this recurrence, right? Because we have this picture in the paper how the model, for example, converges quickly when the next token prediction is easy, but the convergence is much slower if the next token prediction is hard, or if it's a token that's somewhat important for the question to understand the question. right and we just we totally didn't train for this we trained with um we take fixed sequence we take um training sequences and we assign them the same step budget for the entire sequence just because that makes sense for pre-training scale right and this like per token specialization just seemed to emerge with scale that the model actually converges quicker on easy tokens than on hard ones, which was pretty interesting.
24:04And we really didn't know it would work like this before we saw the final model. So elaborate on that. When you refer to per-token specialization, what do you mean by that? So basically, the way the model was pre-trained was not only do we pick a random step for every recurrence. But we also pick the same step count for all the tokens in an example sequence. Because basically for this to be parallelized, the model has to do the same amount of compute at every token. This is trained on sequences, 4 ,000 tokens long. So meaning that random number is at like a batch level or something. Exactly, right?
24:52We're actually saying, okay, this batch gets three steps, this batch gets eight steps, this batch gets 128 steps. But then we look at the test time and we actually see that the model now has this convergence graph that's very different on a per token level. Where on some tokens, it goes like green very quickly, which means like it converges very quickly. And in other ones, it just keeps on doing things for a longer time. And this specialization really just emerged during training. Okay, so let me make sure I'm understanding this. What you're saying is that you've trained it on just these random numbers on a per sequence level, like random number of steps.
25:32But then at test time, the model actually does apply more compute to more difficult tokens in this case. I was a bit more like dancing around saying like, does it use more compute for that? meaning because that's more nuanced it's more nuanced than that oh it's more like in the chart like in the chart in Hippie what we show actually is that we in that chart we actually run the model for the same number of compute right for all steps that's why this is like a rectangular chart and then on that chart we just see how the convergence like how easily like how quickly it's done in every token but to make that chart we actually run all the steps to completion and from that we imply that you could have exited early but the chart really is quadratic and it has a full compute so basically what I'm really saying is it can do that for most of the paper we haven't done this a lot yet maybe this is the most exact way to say it because we have this section about how you could do adaptive exits and you could exit early and that's section 6.3 in the paper and there we try this But that's also where we discovered that this is possible.
26:45And so a lot of the rest of the paper just hasn't had this yet. A lot of the rest of the paper is like, here's a query. I'm going to run the query with 32 steps. Here's a different query. I'm going to run it with 64 steps. So a lot of the paper hasn't caught up to this more adaptive paradigm yet. Where it's also something that we were surprised it would work at all. And which is something where I'm personally very excited. Thinking back to these halting modules of now that this model is trained, we can probably just fine tune these stopping modules into it again, because now the bootstrap problem is solved and now we can just fine tune these and have these exit conditions again that I learned.
27:25But we just haven't gotten to that yet or no one has, I think. I mean, even before you said this, like it was surprising to me that the random stopping worked at all. Like, and why is that better than just like, you know just maximizing um during training i'm maximizing the number of steps right uh yeah that's a good question right because they have been like they have been previous architectures for example do you not sure if you remember something like albert which was a weight shared bird variant so these and there's also some some utmoe models that have this um where basically they have recurrence but the recurrence is always fixed.
28:11It's always like you always do 8 times or you always do 16 times or something. And those also work but interestingly those models don't really generalize to fewer or more steps of compute. They really train like a fixed depth transformer where they lose performance once they go for more steps and they also don't work if they have fewer steps. and with this objective really try to generalize to maybe really an arbitrary number of steps so the model is stable and can really learn more mathematically you can really learn some fixed point or some steady state behavior and that's what this objective is trying to get it to do because otherwise you would maybe you try you put in more computer test time and then it's it would get worse again it would like veer off track and go into some direction it wasn't trained on and so like right because like actually it really is this very non-linear dynamical system that we're training here and we're trying to get it to do what we want it to do and to predict the next token and that's a bit it's a bit tricky to train and if it's for example if it's fixed then it behaves like it's a fixed transformer.
29:32Is there an interpretation of the paper that is searching in latent space for the next token? It's interesting. It's an interesting way to look at it. In some ways, yes. Especially because if you think about it, what actually is happening is that you have some initial state in the recurrence. And the initial state is always random. And then We can also visualize those and often these random states, they lie on like some, for example, if you plot PCA directions, then you will see these random states just somewhere in the middle because they don't have a significant component in any direction. So they sort of like lie in the middle.
30:14And then you see as the model recurs, it keeps on like moving these points away onto some target position that where it wants them at and where then it can predict the next token from them. and we have these kind of funny little spine charts in the end of the paper where you sort of like see this like middle column rising and then you have these trajectories going out at the end. I think this is somewhere in the appendix where you really see that like from like these like from random it then converges to this or finds this point from which to predict the next token. And what's also super exciting for us there is that this is not always just a point.
30:55It doesn't always just converge to a fixed point. Because sometimes you actually also observe the model actually rotating. It has these little orbits that it makes. And it's kind of interesting because then we have these, right? It's still very hard for me to actually visualize this 5 ,280 dimensional space. So I only can show like, right? I can only show these like first 40, first 80 PCA directions or something, right? But in those, then we can see like, okay, what is it doing? It's like often like there's like some trajectories and they all come together and then they do these little circles.
31:27And that's quite interesting, right? Because then like in later tokens, it can pick up on these circles or in these like these orbits that it had done in earlier steps. And it can do some calculation with that. Because we do know that models often implement additional algorithms with these sort of schemes where you have orbits that reinforce each other. And it was quite interesting for us, right? Because it really was an emergent behavior that it would implement some of these problems like this. It wasn't really trained to have these little orbits, but it just emerged from training that this was a good solution to predict the next tokens.
32:07there's a note in the paper that kind of compares this to diffusion models elaborate on that a bit from some perspective this really is a diffusion model I mean maybe this is a bit generic maybe because all recurrent models are or maybe both of them are recurrent models diffusion models are also recurrent models in a different way it's important to separate that from there have also been some very exciting language diffusion models recently but this one's a bit if you think about this as a diffusion model it would be that you pick one token the next token, right? And then you diffuse all the way until you found the next token and then you go one step forward and you diffuse again until you found the next token and you go one step forward you diffuse again until you found the next token this kind of analogy would work here and during development we actually tried to make this a bit more explicit because we also tried to not only start from a noisy state, but also inject noise proportional to how many steps you've done so far, like you would in a diffusion model.
33:18And this didn't really help doing pre-training, so we removed that part again. But it really shows this analogy that this recurrence is not so different from a diffusion model that would diffuse the next token. What is quite different is that the objective, the training objectives are very different. In a diffusion model also this is why discrete diffusion models or language diffusion was all a bit hard is ideally in a diffusion model you would want this process where you have a you have a continuous target and you add some noise to it and then you learn this denoising operation. That's really how large scale image diffusion models are trained.
33:59And in a language that's not so clear because you can try to denoise embeddings but it's like not so nice to do that there. And so the training objective here is still quite different from a diffusion model, but they are both related models in that. Yeah, they both diffuse the next token. Maybe that's the way to say it. Okay. So you mentioned that you started with a hundred million, I think, parameter model as kind of your test model. And then you ultimately scaled that up to a three and a half billion parameter model. did that all happen in one shot or was that an incremental process and maybe talk a little bit about the process?
Read the full transcript
34:43Yeah, what was kind of tough about this was that ultimately this was pretty much a one shot because we had a limited amount of compute on the supercomputer and which we also had in some little segments where you apply for in a queuing system and then you get some 12-hour slot sometime in the year or sometime in the month, right? and you kind of hope that it's running when it's running and um so it really was like we only had these like little 100 million parameter test models where lots of things worked and then we say okay now we're scaling it up and then there's some um there's some first questions like okay what's the data parallelism strategy right what's the model parallelism strategy that we can do there and there was it was a bit harrowing that it was these um it was an amd gpus so we had a bit of issues and so actually the cluster wasn't really there's some like nodes that it wasn't supposed to really to be running deep learning workloads on more than 128 nodes and there was something that like the interconnect between the nodes wasn't quite stable and but we kind of had this allocation for 512 nodes so for 4000 cards which was a bit required because these are the 250X cards, this is the MD equivalent of the A100s.
36:01So these are older generation cards, so you need a bit more of those to get maybe a similar number of flops that you would have with a few or more modern cards. And the cluster wasn't really set up to do this kind of thing. And then we ended up cooking this distributed data parallel pipeline that pretty much we hand-wrote, and then it only worked if you send packages that are exactly 64 megabytes over the interconnect. And that ended up working, and then it was stable. and it works for no other like node count in 512 this training script and it was a bit of a wacky process to actually get there and to do that and then right like ultimately we we just uh threw out the pytorch ddp and wrote in this like funny little uh version of our own that only sends these 64 megabyte packages because that was somehow optimal to like not overwhelm the interconnect on these machines so this was a whole side quest that took us like several months to figure out how to actually run with this allocation that we were having, right?
37:04And we're academics. We're pretty happy about any computer we have and we can make it work. But it would take us some time to make it work. Is the implication that reproducibility is going to be a challenge for anyone that wants to try to do this on their own? I think it's much easier to reproduce on an NVIDIA cluster. I think that's the easy way to say it. But also AMD has been doing a lot of improvements and the new series is a bit better. And so I don't want to say this wouldn't be reproducible on an AMD cluster. But on this particular machine, this script will reproduce it. But I think like some of these problems are really also related to the interconnect on this machine and to the scaling.
37:37But yeah, we've open sourced all the code that ran on this machine, in this cluster, and we have all everything. We also have the model in the open with all the checkpoints. Do you have a sense for what the cost to reproduce would be like on, you know, if you were to reproduce it on, you know, NVIDIA instances in the cloud? I don't actually have it off the top of my head. But basically, we trained on like about 20 segments of 12 hours. So 10 days. 10 days on 4 ,000 cards. Maybe you convert these cards like from 4 ,000 to 2 ,000 A100s to 1 ,000 H100s. Very naively. Maybe 10 days on 1 ,000 H100s.
38:22Not a small. Probably you could get more compute out of the H100s. but as a ballpark I just make like not a little amount of compute um right um and like the model ultimately in the end is also like a it's not a terrible model which actually was surprising to us so to me maybe to most of it because um well like it was like this like a super small team right we have like a few people that know the HPC and then we just kind of built this model right and this is the first this was the only training run and also we had some idea what the data should be to be good and it was just because I picked that and I said okay let's train on this but there was like no computer to train maybe a twin model on a different data mixture or to really figure this out and what data did you use?
39:09so this is like a so really like the first consideration that we had when we when we made the data set was that we wanted something where we could hedge a bit where the model would be good at So we put in like a larger amount of code, a large amount of LaTeX, some chess puzzles, some math, but also a lot of web text, some instruction data. Just try to see like, okay, where would the strength of this model be? And then that's the kind of mix we trained on. Ultimately, it mixes very heavily on code and math, but also a good amount of synthetic data and instruction data. This was pre-trained. It was trained from scratch as opposed to Attune.
39:47yeah but because this this approach is so different it really had to be trained from scratch to do this interestingly there have been a few other papers that um also try to get this this latent reasoning idea right that you could do more there to try to get that from retrofitting fixed depth transformers which also is interesting right and especially for the comp like there's one paper from this out of google and one out of meta it's very enticing for the companies if you already have your large fixed depth consumer can you retrofit this capability into it but for us we really want to see okay can we actually train from the ground up for this model to be recurrent and to do more compute for harder problems and in the abstract you kind of compared the results to a 50 billion parameter model what exactly do you mean there yeah so there it's always important that we write.
40:42So a good way to think about the model is that it actually only has eight layers. So it's like two layers before the recurrence and it has four layers in the recurrence that are being repeated and then two layers at the end that produce the final next token. So this whole thing is just 3.5 B parameters. But now like how do we but actually if we run this model for more steps let's say we run it for 32 steps we actually run it we actually run the middle part much longer right so it would be as if we had run it for something like 50 billion parameters we actually we still have only a 3.5 but the actual calculation is somewhat that we have we have 1.5 billion so like half of them are fixed parameters or 1 billion is fixed parameters and then 1.5 billion is recurrent parameters.
41:39And so that's like in compute, it still takes as much time to get the next answer as if it was a 50 billion parameter model. But it still is a tiny 2.5 billion parameter model. And we compared it like we actually see that in benchmarks, it's quite good to compare like the other open source models like the Olmo series. We actually were very happy that we could catch up to the Olmo v1 models and like v1.5 so that was quite exciting because those models were trained with a much much larger much more capable team i would say right a very good open source team um ultimately this is a bunch of guys in the basement right um so uh we were pretty happy that we could catch up to the v1 version of that model so like where they were like last march we couldn't catch up to the v2 version because they also improve a lot and uh but what's interesting that for them like i think a lot of improvements also were in data so i think that makes us very hopeful that maybe the approach is not fundamentally limited but with better data with more care with a different like a post-training or mid-training schedule it could be better so that was very motivating for us to see those those numbers at the end yeah but it's not a 50 billion parameter model in comp in it's not like it's a llama 70b model it's not nearly as good right because ultimately it's trained for much less compute still is that part of the implication then that as model architectures evolve and incorporate ideas like this the idea of comparing you know models by parameters is going to be less useful i think so but like for us even this was already uh super tough to compare in parameters so the model is actually smaller than all the other ones we compared to but in compute for example like for some other benchmarks uh right how do we compare this if the model is smaller but we spend more compute on getting the answer compared to this larger models that spend less compute because and right the same thing also happens if probably with the other reasoning models right it's like uh we don't like for example the r1 model maybe it's like 64 billion active parameters but do you so do is it a 64 billion parameter active parameter model or is it like a thousand tokens of chain of thought times 64 billion active parameter model it's not so clear how to do this calculation yeah it's become somehow less meaningful what's also pretty interesting is that um wait so so one thing we can do is we could come like convert in pre-training flops to like what model size could we have also trained for with that compute budget exactly right and that comes up that's that comes up to something like a 32 billion parameter model in raw flops so also not small for like a similar token count um but interestingly we couldn't have trained a 32 billion parameter uh normative former on this on this uh node setup actually because what because because the model is so much smaller it's so much more communication efficient because we actually write this actually less state to send to the other devices and so it's kind of a funny thing that with this cluster actually we couldn't have trained a 32 billion parameter model even though that would be flop equivalent and so these compare these comparisons are becoming even harder in the end we just have a table in the paper with different models that are also open source and we say here you can pick your best comparison from the table.
45:02Did you find that the model performed better or worse for in particular domains? Was there any kind of domain affinity for the model? The model is surprisingly good at reasoning. Like it really is like, this is maybe it's just a headline, but actually in the benchmarks, especially on graduate math, not graduate math, not on graduate math, on grade school math, right? Tiny, tiny, big difference. Yeah, right. The GSM8K is a classical grade school math problem. And also on human eval, which is the coding, right? These are nowadays easy coding questions. There, the model is particularly good. And it actually is better than some of the other open source models by the other teams.
45:43And what's there particularly exciting is that, so while we had limited compute, we did train a baseline comparison model where we trained the exact same architecture, the exact same data, the exact same cluster. But we trained it to always have a recurrence of one. and so that would be a fixed depth just comparison right the model has that model has the same number of parameters and we trained it just the same way and we trained that one for 10 000 steps which is about 180 billion parameters and interestingly that fixed depth version got like a gsm 8k score of 1.5 percent or something and at 180 billion tokens the recurrent model already was at 10 to 12 % of JSMEK.
46:27It was like five times better at grade school math, right? Same number of parameters, just a different way of using them. And it was drastically better at grade school math, which was pretty exciting. And right, like still the comparison is hard, but we haven't trained the baseline all the way to the end. So some of these comparisons still are pretty hard because ultimately we have this one data point, right? This is the one model we trained. So as a scientist, I was like, I always, I'm a bit like, I don't want to say like too much here, right? And I'd say this once will always be great, but it was surprisingly good at code and math.
47:05There's a section of the paper where you talk about recurrent depth simplifying LLMs. Elaborate on that point a little bit. Yeah. So what we find pretty interesting is that this model, or what we want to show in this section, is that this is a very natural way to use transformers. So leaving all the thinking and reasoning, all these more lofty motivations aside, we think it's also a practical architecture. And there's several examples in that section about why it's a practical architecture. So we've talked about these adaptive exits already, because this is something that you can also do in a normal transformer.
47:42And people have tried this, people have tried it like if you have 96 layers, maybe on some tokens, you only use the first eight or something. And then they get into all these, okay, how do you fine tune for this? How do you make this happen? And it's a big effort to find in the model for that. And then we show, okay, actually, this different architecture can just do a zero shot. You can just exit. And something similarly, this also holds for things like speculative decoding, for example. I think you also had a guest a few weeks ago about speculative decoding, right? Mm-hmm. I talked about it with Chris Lott at Qualcomm quite a bit.
48:18Yeah. Right. So To recap that slightly, for specular decoding, ideally you need this little draft model. And it kind of needs to be aligned with the large model and it needs to fit and you need to really make that work and do some engineering. And there was a lot of good research being done to make that work. But what's kind of interesting about this recurrent architecture is that just naturally by just running it for one step or four steps, you naturally just have your own draft model in your back pocket. You can pretty much just run the model for a small number of steps to just draft and then you can still verify with 64 steps later and it just happens like that right there's a no fine-tuning necessary no draft model making necessary just it just fits like that right or something related to that is um it's also this idea that you want maybe you'd want to do kv cache uh sharing where you have several layers that's that share the same kv cache so you can push down on the kv cache and in normal transformers that's a bit hard because all the layers have different KNV projection matrices.
49:28So the KV caches don't really match to each other. But this model, because it's just a recurrence, it's always the same parameters are being used. You can just naturally share KV caches between steps if you wanted to. Does the fact that the parameters are the same also have an impact on memory bandwidth during inference? basically yes because so the model is much smaller right you can run it on a smaller chip and you can just run it longer on a small chip that for that it's pretty great we kind of wanted it or maybe not maybe not wanted it but um i mean a different thing you could want is that the entire recurrence fits into like an l1 cache but it's a bit too large for that but if it would then it would be even faster but but right now still you have uh you have about a billion parameters in the recurrence that you keep on moving out of the L2 into the L1.
50:23So there's still some movement there, but overall, you need much less memory on the device to actually store the model, which we think is also pretty exciting for some low-resource applications. And just for logging fewer parameters around, instead of having these massive models where we use our parameter only once, but then we throw it away. it's quite an interesting paradigm that we actually we really try to reuse parameters and
50:58relating to these natural where this application is natural for the LLM I think also there are other ways to use this in interesting ways that we just haven't figured out yet and these are just some examples that we had in this section that I thought were pretty interesting but But it's a bit of a, right? It's a different architecture. It has these different advantages that we aren't maybe so aware of. Even like putting away all the reasoning aside. This might just be more efficient in some applications. Interesting. Interesting. Did you have any interesting, like, qualitative observations about the generations?
51:36Like, are they, you know, fairly standard? Or, you know, did they differ in some qualitative way than what you might expect for similar models? similar in some metric size or something. What was interesting was that, it's also a funny quirk, is that we actually trained the model with instruction data, so with chat templates in the pre-training, which is a, some other people have also done that. But it's a bit more of a rare thing to do. And we kind of did this because we really also didn't have compute for like a longer post-training or cool-down phase. But what it implies is that pretty much the model learns the chat template from the get-go.
52:20And you can just chat with all these intermediate checkpoints already. A fun thing we think about. But you can just chat with the 10 ,000 step chat and the state of the model and see how it's doing and how it's understanding some questions. And you can also chat with the later ones. Because it just inherently understands this chat template format any like post-training. Which is like similar, this is a bit of orthogonal to the recurrence, right? It probably relates more to the data mixture. You know, so this work is focused on efficiency primarily. You also spend a lot of time thinking about model safety.
53:01It strikes me that there are some interesting implications in a model like thinking to itself as opposed to thinking out loud, you know, via text tokens. Have you thought much about that? Yeah, a lot actually.
53:23So I think the first answer like why I think this is not, is maybe why I think this is not worse, which I also talked with colleagues about, right? And there maybe the interesting answer is that I already didn't think that language models think in text. So for me, this was a very small step because if you think about it, if you have a very long chain of thought, the model actually, I mean, of course it takes these text intermediate steps going as text inputs, but actually mostly the model attends to its own internal representations of these previous tokens. But if you think about the attention and all the layers, you are already attending to these very deep representations of these tokens.
54:10And you sort of like already don't know what's in these states. So text is kind of a reflection of some opaque, you know, latent space anyway. And so the thinking isn't really happening in, you know, in ASCII. It's happening in that domain. Yeah, right. The model is like prompted with a new token and it does incorporate that into its like hidden states. but it already thinks based on its previous hidden states. And with this model, maybe this is just a bit more explicit.
54:47But what I also think is like, well, like why I think this is interesting is that I do think it's a question of like, how do we oversee these models? And to some extent, I also want some of these be, some of these like developments to happen in the open and happen as open source. And this was part of the motivation also why we open source this particular way of doing it. because ultimately, like, I also think it's a very natural way of doing this. Like, I think if we wouldn't have done it, it would have happened. And so for me, this is also a way to say, okay, here is an open source version of this model.
55:22Here's how it works. Let's try to understand it. And that's also what I'm really interested in right now is trying to understand it now that we have it. And is that the kind of focus for you and the group going forward? Is it to evolve the model or specific experiments or incremental directions that you're heading to? Or is it more like we have this now, let's try to understand it better? Or, you know, probably all of the above to some degree. But yeah, it's really both of the above. Actually, the most immediate thing is that the model actually is not even cooled down, right? Just because we actually don't have any compute anymore.
56:01That's the very real answer why it is not cooled down. What does that mean? For example, like for the later Ulmo models, we see that the models improve a lot with this mid-training phase, where the model is slowly cooled down from its peak learning rate to a zero learning rate using a very high-quality data mix. So for me, like on the first part, I think on the agenda is to try some data mixes, for example, maybe tailored more towards code, more towards math, and you'll see what these cool-down behaviors will look like for this kind of model. And that's already interesting, right? Because you can also play around with the parameters.
56:33For example, you could imagine that you could cool down with more steps than you did during pre-training. So you could incentivize to think deeper during the cool-down phase. But there's also some other questions that we have on the team of like, how would you even post-train this kind of model? How would you do, you know, now we become all the way back around. Yeah, I was going to ask that. What does fine-tuning mean here? Yeah, right. Or now that we've seen R1, like what's the equivalent here? how would your reinforcement learn for this kind of model architecture for this kind of paradigm right and these are questions that are in the room aside from questions of how do we what is it doing internally right how can we understand how can we probe what this model is doing internally yeah so these are all things we're thinking about but it's all pretty exciting it's very cool stuff and I'm glad we're able to connect and chat about it a bit yeah I'm also happy to field field questions and anyone who has like my email or anything like this is like I've really been thinking a lot about this model and about you know how to do it how to train it and also trials and tribulations so I'm always happy to answer more questions about this awesome awesome thanks so much Janice yeah thank you for your time
57:59Thank you.
From the publisher
Today, we're joined by Jonas Geiping, research group leader at Ellis Institute and the Max Planck Institute for Intelligent Systems to discuss his recent paper, “Scaling up Test-Time Compute with Latent Reasoning: A Recurrent Depth Approach.” This paper proposes a novel language model architecture which uses recurrent depth to enable “thinking in latent space.” We dig into “internal reasoning” versus “verbalized reasoning”—analogous to non-verbalized and verbalized thinking in humans, and discuss how the model searches in latent space to predict the next token and dynamically allocates more compute based on token difficulty. We also explore how the recurrent depth architecture simplifies LLMs, the parallels to diffusion models, the model's performance on reasoning tasks, the challenges of comparing models with varying compute budgets, and architectural advantages such as zero-shot adaptive exits and natural speculative decoding.
The complete show notes for this episode can be found at https://twimlai.com/go/723.




