Teaching Large Language Models to Reason with Reinforcement Learning with Alex Havrilla - #680

16 Apr 2024 · 46 min

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

Podcast Summary: Teaching Large Language Models to Reason with Reinforcement Learning with Alex Havrilla - #680

Podcast Overview Title: The TWIML AI Podcast Host: Sam Charrington Guest: Alex Havrilla, PhD student at Georgia Tech Episode Focus: The application of reinforcement learning (RL) to improve reasoning capabilities in large language models (LLMs).

Key Themes and Concepts

Introduction to Alex Havrilla

  • Background: Third-year PhD student at Georgia Tech, involved in neural network learning theory and open-source LLM research.
  • Research Interests: Focus on RL fine-tuning methods and enhancing reasoning in LLMs.

Research Focus

  • Reinforcement Learning for Reasoning:
  • Applying RL, particularly automated forms (without human feedback), to enhance reasoning capabilities of LLMs.
  • Development of the TRLX library for conducting RL fine-tuning experiments.

Differences in RL Methods

  • RLHF vs. Broader RL Application:
  • Discussion on how RLHF (Reinforcement Learning from Human Feedback) improves LLMs, but emphasizes that automated RL can also yield significant results, especially for reasoning tasks.
  • Notably, RLHF is seen as pivotal, but automated RL offers potential for scalability.

Data Efficiency in RL

  • Sample Efficiency:
  • Traditional RL methods are often data inefficient; however, fine-tuning pre-trained LLMs demonstrates immediate improvement with much lower sample requirements (often fewer than 100,000 samples).

Applications of RL in LLMs

  • Potential Uses:
  • Aligning LLM interaction with human expectations (RLHF).
  • Allowing LLMs to utilize tools (e.g., web browsing) effectively, enhancing their interactive abilities through RL strategies.

Achievements in the Research

  • Sample Complexity Findings:
  • Exploration of different RL algorithms (PPO and expert iteration) showed surprisingly similar sample complexity in practice despite differing theoretical expectations.
  • Investigating methods to maintain output diversity and solution exploration.

Importance of Diverse Outputs

  • Role of Diversity in Model Responses:
  • Diverse outputs lead to enhanced exploration and improved learning, particularly in reasoning tasks.
  • Over-fitting during supervised training could harm diversity, leading to poorer performance in RL settings.

Future Directions and Insights

  • Potential Applications for LLMs like Llama 3:
  • Emphasis on exploration as a critical factor for improving reasoning capabilities.
  • Suggestion to benchmark various algorithms for generating synthetic data and improving model training.

Challenges with Noise in Training

  • Static vs. Dynamic Noise:
  • Static noise (local errors) can be tolerated to a certain extent in training data without severe impacts on model performance.
  • Dynamic noise (errors affecting the entire sequence of reasoning) is more detrimental and leads to reduced learning capabilities.

The Advanced Reasoning Benchmark (ARB)

  • Comparison with Other Benchmarks:
  • ARB aims to provide challenging questions to assess LLM reasoning, moving beyond just numerical final answers to include symbolic reasoning and proofs.

Key Takeaways

  • Reinforcement Learning is Promising: The potential for RL to enhance reasoning in LLMs is significant, and ongoing research is key.
  • Diversity in Outputs is Crucial: Ensuring LLMs generate diverse outputs supports better exploration and learning.
  • Future Integration of Systems: Combining LLMs with external systems can enhance reasoning capabilities, driving towards superhuman performance.
  • Understanding Noise Effects: Research on how noise in the training data affects model performance is essential for refining training methodologies.

Conclusion This episode emphasizes the transformative potential of reinforcement learning in enhancing LLMs' reasoning capabilities, highlighting the importance of exploration, data diversity, and the need for rigorous benchmarking in future models. Alex Havrilla's insights into the interplay between RL and LLMs provide a valuable perspective for researchers and practitioners in the field of AI.

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

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:05All right, everyone, welcome to another episode of the TwiML AI podcast. I am your host, Sam Charrington, and today I'm excited to be here with Alex Havrilla. Alex is a PhD student at Georgia Tech. Before we get going, be sure to take a moment to hit that subscribe button wherever you're listening to today's show. Alex, welcome to the podcast. Thanks for having me. I'm excited to be here. I'm excited to dig into our conversation as well. We're going to be talking about your research, which centers around the idea of applying reinforcement learning to reasoning and some of the downstream problems you ran into in the process.

0:37But before we dive into your specific research, why don't you introduce yourself a little bit to our audience, share how you came to work in the field, and we'll start there. Yeah, so like you mentioned, I am a third-year PhD student at Georgia Tech. Technically, I'm in the mathematics department. My advisor does neural network learning theory. But like two years ago, I got into open-source LLM research through the open-source community, and since then have been very interested in RL fine-tuning for LLMs, specifically with a focus on trying to improve their reasoning capability. Tell us a little bit about how you think about your research broadly.

1:16Broadly speaking, I am kind of dividing my intention into roughly two parts. So the first part is working on this more theoretical aspect, like I mentioned. So neural network learning theory, trying to, you know, say theoretical statements about given a certain number of training examples, how quickly can your model attain some minimum generalization error, how large your networks need to be in theory to, you know, approximate some target function. And then on the more applied LLM side, my research is kind of a little more broad. I focused on RL fine-tuning. I was responsible with some other very talented people for the development of the TRLX library for open source RLHF.

1:58And that allowed us to do large-scale LLM RL fine-tuning experiments for RLHF, which kind of paved the way for doing large-scale experiments with reasoning, which was always my original motivation. And when you talk about and think about RL for reasoning, how much of that is RLHF versus the more broad application of RL-related methods to improving the reasoning capabilities of language models? There's no human feedback. So this is entirely automated. You can get whatever signals you're interested in from some ground truth label. So maybe it's very distant human supervision, but in the training loop, there's no supervision.

2:36But certainly I think human feedback is also a very valuable way of improving LLM reasoning with RL. So one great example of that is this OpenAI Let's Verify step-by-step paper, which trains this process-based reward model to give very granular step-level feedback to LLM solving math problems. And lingering on RLHF for a moment, you know, when you think about that as kind of the composition of RL and HF, you often think that we get wrapped up in RL and really the heavy lifting is done by the HF. Do you agree, disagree? How do you think about it? So I think it's for sure true that the majority of the improvements that you see in today's LLMs are due to the human feedback parts.

3:18I mean, I think that was evident even if you look as far back as the Instructs GPT paper. Like they had these comparisons where they have this like kind of intermediate model, which they called the FeedMe model. This is TextAVinci 2. And that model was already like miles better than GPT-3. I mean, it was just much, much better. And then like once you do the RLHF on top of that, like the whole PPO pipeline thing, it improved somewhat. Like there was a 10 % preference win rate improvement over that, you know, FeedMe SFT baseline. But this is, you know, nothing compared to the improvement you saw over the original base model.

3:52So I think, yeah, it's definitely true that the human feedback part is doing a lot of the thing. But on the other hand, I think, you know, the hope that many people have for RL, myself included, is that, you know, if we really want to, you know, start training these superhuman level systems, you know, the amount that we're going to be able to provide high quality supervision is just going to naturally decreases. The complexity of the tasks we attempt to solve with these systems gets bigger and bigger. Continuing along the RL vector, one of the historical challenges with reinforcement learning is its data inefficiency.

4:28Do you see us solving that by leveraging kind of unique properties of language models to make RL more data efficient and useful in this particular context? or is it more about, you know, tracking improvement in general RL research and kind of pulling that into application with language models? So like you said, more classical RL, yeah, extremely sample inefficient. Like you might have to wait around for like, you know, hundreds of thousands of rollouts before you see like any emergent behavior at all, before your agent learns anything. Not the case with RL for LLMs. So like when you start, you know, fine tuning these pre-trained LLMs, they start reliably improving like basically immediately.

5:16And this is, I think, because of like, there are a couple differences here. But the main difference is, you know, like when you're fine tuning this LLM, you're starting, you know, either from some supervised fine tune checkpoint or from, you know, some like pre-trained checkpoint, which already has a pretty good sense of what it's trying to do. Right. So like there's a very strong, like warm starting bias, which you don't have in classical RL. And that was like really good for sample efficiency, because you don't have to wait around for the model to kind of like, you know, make random moves and then see if those random moves are useful.

5:45You can just, you know, immediately start generating in the context of reasoning, you know, you can immediately start generating potential candidate solutions, some percentage of which will be correct. And, you know, you can get a pretty good reward signal from that immediately. And so that means sample efficiency is pretty good too. Like actually, in most of the experiments we did, we didn't require more than 100 ,000 samples most of the time to like, you know, really fully converge to the max performance that we see. Yeah, maybe with that in mind, kind of frame out, you know, where we are with RL applied to LLMs.

6:20What are the, you're applying it to reasoning? Are there other ways that RL is promising in the application to LLMs? And kind of where are we broadly speaking? Yeah, so of course, there's RLHF, like we talked about, super useful for, you know, aligning the interactability of the model with how humans expect to interact with it. So there's that. I think also there's a lot of, you know, work going to be being done and will continue to be done integrating LLMs with, you know, various tool usage and allowing them to, you know, act as web agents, for example, browsing the internet, these kinds of things.

7:01And there will definitely be RL components there, because that's a very interactive process, right? Like when you're learning a new tool for the first time, you have to figure out what are the scenarios in which should I use this tool? What are the preconditions for using this tool? And anytime you have an interactive process that is very naturally RL focused. Given the cost of training large language models and their complexity, is there an online component that RL makes possible for LLMs? How do we implement online learning in LLMs? So first of all, online learning is still a very difficult problem.

7:38So like this idea of continual learning or updating yourself as you go, still quite difficult. I mean, I think there are some naive solutions like gather some data, fine tune the LLM on the data, and maybe those work okay, but it's kind of hacky, right? It's not really the solution we're looking for. But I think you bring up another good point, which is, and this is one of the other, like one of the main things we were investigating in the paper is these different paradigms of doing RL training. So like there's this online on policy paradigm where you have some agent in an environment, right? And you're rolling out the agent.

8:10And then you're training the like policy, you're updating the policy using the like data, which has just been collected by the agents. So that's online and on policy. The agent's doing exploration and you're updating the agent using data, which is generated. but there are kind of like more off policy versions of this where you can generate lots of data and then you can the agent you're training and updating is like updated on a kind of old or stale data so it's like put concrete names behind this proximal policy optimization ppo is an example of like a you know a mostly online on policy algorithm and that's one of the algorithms we looked at in the paper.

8:49Expert iteration, you have a question and you sample the model many times on the question for many different solutions. And then you fine tune the model on the correct solutions. That is much more off policy because you're generating lots of data and then you're fine tuning the model on trajectories, which may come from different points in like the training cycle. So and usually like when you're doing this type of training in the classical RL setting on policy algorithms such as PPOR, they're not more sample efficient technically, but like you need less iterations of like training to fully converge because you're exploring this complex state space and somehow the updates you make to the policy are like the most useful ones you're making in the context of solving the problem.

9:37And off policy is somehow like converges slower in comparison in classical RL. But this is not at all what we found in when we were comparing these methodologies in RL fine tuning for LLMs. So when you look at the sample complexity or the number of samples you need to reach a certain level of performance for PPO versus expeter duration, it's roughly the same, to be honest, despite the fact that expeter duration is much more off policy and potentially generating less useful data. which I think a priori is a bit of a surprise because this is, again, not the case at all in classical RL. So these two methods have roughly the same sample complexity.

10:15And like I mentioned before, the overall sample complexity you need for these models to converge is pretty small, like only 100 ,000 samples. So somehow this is an indication that definitely the models are learning something and are uncovering interesting behavior, But somehow like, you know, they're not really accessing, you know, they're not somehow diversely exploring in the way that we might want them to. Because otherwise, if they weren't accessing, you know, like new types of solutions and new modes of behavior, you probably wouldn't see such like fast conversions and like such similar sample complexities between these methods.

10:52Let's maybe punch into the paper so that we can talk about some of the broad RL concepts and their applicability kind of more concretely. So the paper is, again, teaching large language models to reason with reinforcement learning. Talk a little bit about your motivations with the paper. Yeah. So I already kind of touched on it. But, you know, in the RLHF space, there are many different algorithms now that are popular, right? So at the beginning, there's PPO and PPO with SOTA. And it was unclear whether you could use other algorithms. But now, you know, I'm like, April 2024, it's clear that, you know, not only are there many competitors, but they're basically just as good, right?

11:35So there was a recent paper from Cohere, which shows you that reinforce is basically as good as PPO in this context. DPO is obviously very popular now and does no exploration at all, which is also maybe another sign that something is not quite right here. But it's basically just as competitive as PPO, which does lots of exploration. But when I was investigating these problems, this was not clear. It was not clear how do these algorithms compare? like what are the benefits of using one RL algorithm versus another for doing this fine tuning. And so this was really our goal, like we're gonna take all these algorithms and we're gonna take a bunch of different ideas in RL and we're just gonna see what works best when you're fine tuning RL.

12:18With the context here, specifically human feedback or something else? This was not RLHF. So in the context of reasoning, usually you have a question and then you have an answer. And that answer has like a goal, final answer you're trying to reach. So like, you know, maybe it's some math word problem and you're trying to get some integer solution to the math word problem. And then, you know, we can give a reward in the most basic sense. We can provide a reward feedback based on did you get this correct final answer or did you not? And that's how you can start to do the RL. What's the relationship between the policy and the language model itself?

12:55Okay, that's another good question. So in this case, the policy is exactly the language model. So the language model generates every single action and then, you know, is supposed to come to the correct final answer. But I think this is a good question because there are situations in which, you know, the policy is not necessarily just the model. So, for example, you can imagine combining a calculator with an LLM to solve math problems. We didn't have to do it here because we were working with LAMA 2 based models. And LAMA is pretty good at arithmetic or at least the types of arithmetic we were doing.

13:30There was no need for a calculator. But you can imagine combining a basic tool like a calculator with the LLM, and that forms a more generalized policy, which is the combination of an LLM and some tool. So yeah, just to kind of drive home the point that the policy doesn't necessarily have to be the LLM. If the policy is the language model, what does it mean to optimize the language model? Are you optimizing, or what are your levers? Are they things like temperature and things that go into a prompt? I'm imagining that they're not the model itself. In our case, it is the model. So like we are literally, we are literally like changing the weights of the model depending on, you know, the type of reward that, you know, is getting back.

14:12Like, did you get this question correct? If yes, you know, like change your gradients in a way to, you know, adjust for this. If you got it wrong, okay, maybe change your gradients in a different way. Got it. Yeah. I was imagining that the model was embedded in some larger system and you were tweaking something about the larger system as opposed to the model itself. I see. Yeah, we're directly fine tuning the model. Got it. And so you alluded to one of the big results that you saw that the algorithms all performed largely the same. Is there a next level of nuance or detail there? Yeah, yeah, I think so.

14:51So let me expand on that. So we went in expecting, you know, PPO to do quite well, like get the best performance and have pretty good sample efficiency. Because like I said, you know, it's more on policy, it should be more, you know, efficient in some sense, it should converge faster. But yeah, I mean, this is simply not what we saw. So like I said, expert iteration, for example, where you simply fine-tune on the correct answers performs roughly the same. In our setup, technically, it was somewhat less sample efficient, but we did some ablations with Stem and Strate. Okay, so to be very concrete, the way we implemented expo iteration is, you know, we sampled the model like 96 times per question or something, and then took all the answers which were correct and fine-tuned the model on those and iterated that.

15:42But then we did some ablations. Correct in this case is determined by human feedback. Right. Yeah. Because like for each answer, you have like a ground truth, correct final answer. And that has been labeled by a human. In the training, you're just, you know, doing text comparison or something automated to make sure that the correct answer is there. Exactly. But you do have the correct answer because they're ultimately math. You know, it sounds like were they all simple arithmetic or is that just one category of reasoning problems that you. Yeah, they were mostly simple arithmetic. We also considered some common sense question answering, but the majority of it was arithmetic, like GSMAK style questions.

16:22But just to finish describing the setup very concretely, so for expert iteration, we had 96 samples per question, and then we fine-tuned on all the correct answers. but it turns out like 96 questions 96 answers per question is kind of overkill like you can reduce that to four answers per question get nearly the same performance as you had before like after each iteration and then you can match the sample complexity of PPO and then you still beat PPO's performance so yeah and then we investigated some other setups as well so we investigated in the offline case where you don't do any exploration.

17:03There's this idea of conditional fine-tuning where you can label each step as correct or incorrect using good or bad tokens. And so we thought, why was this the case? Why is everything performing roughly the same? Because we also tried some other ideas as well. So there are these ideas of training of reward models in the literature. So this idea of an outcome-based reward model where you can predict the correctness of the final answer. And somehow this might smooth the type of reward that the model is getting, because previously you only get a reward of plus one if you got the correct final answer and you got a reward of zero otherwise.

17:44But maybe you're really close, but just slightly off in some way. So hopefully the reward model can tell you that. So we experimented with giving that as a reward, both at the very last token and also kind of interspersed throughout. So it's this idea of a sparse reward versus a dense reward. And again, we found in both cases that these mechanisms improve the sample complexity of the learning algorithm. So if I give you an ORM for a sparse reward or a dense reward, you need slightly less samples to converge. But the peak performance which you get to is about the same or slightly worse than if you just used the ground truth, exact match, text match reward.

18:27And like, you know, we were thinking, why is this the case? You know, like, why is it that even when we give this, you know, better reward, the algorithms all must roughly converge to the same thing. And so our main hypothesis was that, okay, if these algorithms, especially PPO and exporter duration, are all producing roughly the same data, if they're all producing roughly the same exploration of the types of answers which they discover in the exploration process, that would kind of explain this phenomenon, right? Because regardless of the algorithm, if I give you more or less the same data, then they're probably going to perform roughly the same in this context, at least.

19:12So to verify that, we had to look at some measures of what is the type of output each of these methodologies or algorithms are producing. And in particular, we were interested in two metrics. So we were interested in looking at the diversity of the solutions which each algorithm generated. So what is the unique number of solutions for a given problem where you judge two solutions the same if they have the same order of operations or something like that, arithmetic operations. So that was one metric we were looking at, output diversity. And then we were also interested in, if you look at the pass at 96 metric of these trained models, which is if I sample the model 96 times on each test question and then check if I ever get the same test, if I ever get the correct answer, that's another metric I can use to evaluate these models' performance.

20:12And that's also some measure of output diversity because, you know, there's this, roughly speaking, the idea of a model being more diverse means that if you sample it many times, it'll have like a broader set, you know, broader coverage of the types of solutions. And its odds of getting the correct solution are, you know, better than the less diverse model. And so their pass at 96 score will be higher. So we were comparing the pass at 96 score of these models, and we found some very interesting things. so for example if you do supervised fine tuning and you train for like two epochs on some like Golan ground truth data your pass at 96 score will be pretty good so like you can get a pass at 96 score of like 80 % on Llama 2 7b if you do this but then the interesting thing is if you train for longer than that so if you train for more than two epochs Your POSIT 96 score will start to go down and suddenly it seems like your model is producing less diverse solutions and the exploration it's engaging in is worse.

21:19So like somehow supervised fine tuning is not like too much of it is damaging model diversity. I'm not sure I'm clear on the role and value of diverse responses for simple arithmetic problems. So I agree. Like if you're trying to solve a math word problem on the face of it, like, why do you care if you are able to solve it in like multiple ways? Right. So first of all, I want to clarify that even in these simple math word problems, there are a surprisingly diverse number of ways to solve them. So like on average, I can tell you that the 7MB model discovered like five different unique solutions per problem, which is kind of surprising, right?

22:00But like, you know, already pretty diverse. So like, why is this useful? I'm making the claim, it depends on your use case. So like if you're interacting with ChatGPT, for example, and you know, you're just asking ChatGPT, ChatGPT solved me this word problem. Then I agree, you don't care about output diversity. you just only care about chat gpt getting the right answer and if they can do that reliably in the first try that's great okay and so yeah you don't care but like let's say you're interested in like maybe you have like some two two different types of like difficulty of questions right so like you have an easy set of math problems and you have a hard set of math word problems maybe I can convince you that if you can generate lots of diverse types of solutions to the easy math word problems somehow this makes it easier to also generalize to solving the harder math word problems a model with a greater agility with regard to reasoning can probably do better with the more complex problems whereas for the easier problems it might be more like you know memory or retrieval or something like that yeah the reason I think we care about this very concretely in the RL context is because the ability of the, you know, agents that you're fine tuning to produce diverse solutions directly impacts the quality of the exploration it's doing as it's being trained, right?

23:25Like if I have a LLM, which is being fine tuned with RL and it's like, you know, super overfit, has no diversity, it's like just not going to be able to discover that much and like not going to be able to learn that much beyond like what it already knows, right? Because it's simply not generating surprising or unexpected solutions. Whereas if you have an LLM which is not overfit and is able to generate more diverse, interesting solutions, then you're going to start to see the more complicated type of generalization that you're interested in. What's the relationship between exploration and diversity responses and temperature, which I mentioned earlier with LLMs?

24:05temperature definitely will have an effect on exploration uh and like model output diversity um in general in our experiments we fix so it depends on the setup we have two broad setups one we supervise fine-tune the model beforehand and then do rl and the other we fine-tune the model from the pre-chain checkpoint and then do rl um in the case of supervised fine-tuning setup you can afford a slightly higher temperature of like 0.7 because the model already has a decent idea of what it has to do. When you start from scratch, you need to have a lower temperature of like 0.1, 0.2 because the model just doesn't have a good sense of like, you know, what actually constitutes a correct solution.

24:50It's still learning the syntax, that kind of thing. And so you can't really afford to be slightly off. You really need to like make the best of all like the valid examples you have. But what you can do is as you go through training and the like the scratch, like starting from scratch case, you can start to anneal the temperature so that you're producing like higher temperature solutions as the model is, you know, better learning what constitutes a valid solution. Are you tuning just the last layer or kind of how deep are you tuning and that kind of thing? They're all different things you could play with there, but they're not fully specified, I don't think, or didn't think by the RL algorithm itself.

25:30Ah, yeah, yeah, I see. I misunderstood. So, yeah, from that sense, like what layers of the model are you fine tuning? That's all the same between different algorithms. Yeah, yeah. but like how you actually compute the gradients and like the loss which is used that's algorithm specific what's the next big takeaway yeah so another interesting thing is so I mentioned like this phenomena of and this was observed in the original GSMAK paper actually so I don't want to claim credit for it but there was this phenomena where if you you know have a model which you're supervised fine tuning if you do it for like two epochs it has pretty good diversity but then diversity pretty quickly decays after that so like Hopefully I've convinced you at this point that it's in your interest to maintain model diversity if you want good RL.

26:16So the nice thing that happens when you do any kind of RL fine tuning, PPO, expediteration, anything really, is that that pass at 96 parameter kind of maintains, like it converges pretty early in training, but at least it doesn't decay. So like maybe it converges to 80 % and then it stays at 80 % for the majority of the RL training versus an SFT at like, it's this unimodal thing where it like increases and then decreases. So somehow like RL fine tuning is able to reinforce and preserve like the type of diversity that the model has like picked up on and started to learn. What you're really doing during RL fine tuning is you're expanding the kind of like this database of training data that you have.

26:58And just like training the model on more diverse data allows it to maintain more diverse outputs is what I think is really going on there. But that's kind of like one nice byproduct of the RL fine tuning is that you maintain this level of diversity you wouldn't otherwise get with supervised fine tuning. There's this interesting literature in classical RL, which is autocurricular. So like if you have some agents on a bunch of different environments, some of the environments will be harder than other environments. And so you want to like, you know, train them easy and increase the difficulty over time.

27:29Yeah. Yeah. That has worked pretty well for classical RL. So we tried a couple different algorithms like that for, in our case, with the LLM fine-tuning. So to be concrete, we tried PLR, which is Prioritized Level Replay, which is one algorithm. And we also tried this idea called backtracking, where you start the LLM at a solution, which is almost done, but not quite. and so the LLM only has to fill in like the last one or two steps. And then when the LLM demonstrates, it can fill in that last one or two steps, it backs itself up slightly. And now it has to fill in like the last three steps. If it can do that, the last four, et cetera.

Read the full transcript

28:11So you're kind of like backtracking to like early and earlier steps in the solution. You can see that as like being easier than solving everything from scratch, right? So we tried these kinds of ideas. And again, I mean, there were like maybe some small benefits, but they weren't really inducing the type of generalization improvement that we were interested in. The model is kind of already going to, you're going to sample it and it's going to produce solutions in some specific way. And all of these algorithms that we tried are not really doing a good job of significantly changing the sampling distribution of the model.

28:52And so that's why their performance is all roughly the same. uh kind of given what you learned in this research if you were you know working on llama 3 for example and wanted to ensure that it had you know better reasoning capabilities like is there an obvious way that you would kind of integrate in the results of this research yeah so i think at my point my opinion is exploration is the fundamental problem here i think it's clear that there are a couple of things which clearly do work well at improving like the quality of model outputs. So if you somehow have really high quality data coming from another source, you know, like humans, for example, or GPT-4, that clearly works very well.

29:35You can fine tune a model on those and the type and like of solutions that will produce is like significantly changed by this. I'm sure Meta is already doing this, right? Like, of course, they're getting humans to annotate high quality data. If you want to take the extra step, what you really want to do is you want to start looking into comparing and benchmarking how different algorithms impact the quality and diversity of synthetic data that you're generating. So people right now are very interested in synthetic data generation, but there isn't that much investigation yet into how does the quality of the data generated by one algorithm compare to the quality generated by another algorithm?

30:18How does this compare as a function of fixed compute, fixed number of samples? Also, how is the diversity impacted? I think really you want to have a rigorous benchmark which tells you, if I have these different ways of generating synthetic data, how do I get the best bang for my buck when I'm comparing some target's level of quality diversity in the output. You also have recently been involved in a benchmarking effort, a data set called ARB, Advanced Reasoning Benchmark. What's the relationship or contrast between that and GSMA-K, which you referred to? Yeah, so ARB, ARB, is just impossibly difficult.

31:03Like, so difficult that certainly I would not be able to solve most of the questions without better reference solution. And this is really the direction most benchmarks are going nowadays because LLMs really just are so capable. You really like, so like a popular benchmark right now is graduates, it's GPQA, graduate something questions, which is similar to ARB, I think like a bit different in the sense that one of the things we were really going for in ARB is we want like really high quality, really difficult questions, but we also don't just want questions which have like a numerical final answer, which you can check that way, right?

31:41So like we certainly have those in the benchmark. So like really difficult questions that you have to get to correct numerical final answer. But we also have some questions which investigate like the ability of the model to do symbolic reasoning. So like maybe the final answer is some polynomial or something, which you need to get to. Like a lot of, you know, more advanced math questions and physics questions are like this, where like, you know, your answer, your final answer is not numerical. It's also symbolic. And, you know, this is obviously an important type of question to be able to answer.

32:12Also proofs. Like we have a couple high-level proof questions in there that the model needs to be able to, you know, solve and like ready convincing proof for. GPT4 is like pretty good at doing arithmetic, right? Like I think a lot of people have remarked on this. If you ask it to like multiply two five-digit numbers, it probably will be able to. But then if you are like asking GPT-4 to add like two very simple numbers, like asking it to add, you know, 10 and 11 or something in the middle of a very long computation, like, you know, it's integrating stuff and it's like summing some power series.

32:44It's like doing lots of complicated stuff. GPT-4 will tend to mess up very simple like arithmetic operations, which it was completely fine and like able to do in like a clean context, which is kind of funny, right? Like this happens to me all the time. I'm doing some hard math problem and very high-level abstract stuff. But then when I actually have to go and simplify, I suddenly freeze. And I'm like, oh, wait, I'm making mistakes here. I have to go back and check my work. So somehow it's funny that GPT-4 exhibits this same, I would say, characteristic that humans do when they start to learn higher-level mathematics, this kind of thing.

33:22I don't know why exactly that is. it might you know in some way reflect just the difficulty of reasoning over long context actually I'm sure it does I'm sure that's part of it but also I wonder if it in some way reflects the training data itself like maybe you know the training data you have for these types of more long complex questions is filled with like arithmetic errors this kind of thing which you just don't detect because like you know when you're looking at a graduate level question you don't actually care too much about the arithmetic itself, right? Like you're, you're more interested in the more complex, sophisticated reasoning.

33:59Attention is obviously a super overloaded word, but it strikes me that there's maybe something in what it's paying attention to as it's what it, what it is, how it's conceptualized the goal, you know, from a distribution perspective of the output as it's generating the tokens, uh, and it maybe gets distracted by the bigger picture and misses the smaller things or something like that? It's distracted by, you know, this integral and it suddenly can't add like six and five or something like that. You expect it to be a lot more consistent. If it knows how to add six and five, it's going to know how to add six and five.

34:31Right. Yeah. It's just fascinating that it's not like that. So. How has it changed the way you think broadly about reasoning, machine reasoning, LLM reasoning, and the future of reasoning? We didn't get to touch on noisy algorithmic chain of thought. If I'm teaching LLMs to do addition or like multiplication, like very simple algorithmic tasks, how does their performance depend on like, if I start injecting, you know, little tidbits of noise here and there. So like, maybe I flip this digit from five to seven, or like I delete this line. It sounds a little bit like kind of the approaches that are often taken for like interpretability research.

35:10Like let's fiddle with the input data and make it a little noisy and see, you know, how that changes things. Yeah. We're specifically looking at like noise effects training for chain of thought. And like in preparation for this work, I spent a lot of time thinking about like, okay, why is chain of thought useful to begin with, right? Like, it's also clearly very useful to humans, like humans, you know, had spoken language for a long time before, you know, like writing things down became useful. And like, suddenly, I think like, it's clear that once humans started writing things down, this was a significant shift in like, you know, the trajectory of the human race or whatever.

35:44I'm no authority on history. So this is just my I take. But there's clearly like, you know, a huge power to writing things down. And I think like one way of interpreting it is it's like a way of maintaining state. You know, like if I am doing this long, complicated computation, I'm like adding up to seven digit numbers. I have to do all of that in my head and I have to maintain the state and like individual neurons on my head. Right. Whereas if I'm writing things down and doing step-by-step chain of thought reasoning, I'm able to offload most of that state onto the sheet of paper, right? And then I can just refer back to the sheet of paper for like, you know, certain subsets of information that I need.

36:27And that's hugely useful. Like that just, you know, simplifies significantly the type of computation that needs to be done normally. Certainly for humans. And I think also for LLMs as well. Like it just, it literally makes the next token prediction task much, much easier. Chain of thought has proven to be useful and based on chains of thought that have naturally emerged in the training data, but there's some way to maybe take advantage of or kind of design for the way LLMs might think about chain of thought and like train on more specific chains or something. Getting LLMs to do like refinement and this kind of thing, right?

37:13Like, you know, you have some solution and you try to fix the solution. And often when humans are doing refinement, like there's this whole intermediate thought process, right? Like you have some sequence of thoughts. You say, oh, there's an error here. You try something and that doesn't work. You try another thing. But that's not written down anywhere, right? So like oftentimes we get to the correct final answer, but we don't show our work. And so a lot of that work doesn't show up in the LLM training data. But I think putting those intermediate steps in the LLM training data will be super helpful.

37:44I think that's something probably pretty much everyone agrees on and most people are doing. Our key questions are, A, how do you characterize this type of noise in chain of thought data? Are there different types of noise which impacts the model differently during training? and then B, how can you quantitatively model the impact of this different levels of noise on like, you know, trained and model performance? If you're training, you know, GPT-N, you might be able to profile some new data set that you have access to, characterize this noise and then understand the impact that it's going to have on the model that's produced.

38:21Yeah, that's exactly the idea. Okay. Like, I think it's important to be able to identify these factors, right, and say like, yeah. So now to explain what we found. So we trained a small model. This was a relatively small-scale study. We trained a small model, Pythia, 410 million parameters, on just sequences of text produced by simple functions on integers. So we did arithmetic. If you have a list and you need to find the median in the list, we did that. We did sorting. We did a bunch of GCD, greatest common divisor. We did a bunch of different algorithms, which produced like chains of thought for, we call them algorithmic chains of thought because they're algorithms.

39:05And then we trained, we like trained models on them. And first of all, like chain of thought makes the task much easier. Like if I am training even like, you know, a small GPT-2 model to do arithmetic, it's much easier if you write out the chain of thought for the model versus training from scratch, like just going from like X plus Y to Z. That's much easier. So then once we confirm that, we characterize different types of noise. So we came up with two different types mainly. So there's static noise, which is kind of noise, which like maybe makes local changes to each chain of thought. So like maybe you flip a digit or maybe you delete a line, but that doesn't change anything like earlier or later in the chain of thought.

39:50It's just that specific line. But then we also have a kind of more pernicious more dangerous type of noise called dynamic noise. And what dynamic noise does is when you're actually executing, like when you're actually getting the chain of thought, which you then feed into the LLM, you like inject some noise there in the generation process. So like maybe at a step where I'm adding two numbers, I flip one of the digits that I'm adding, producing the wrong answer. And then that impacts all your calculations downstream. So like suddenly dynamic noise has a much larger, less localized impact on the entire solution versus static noise.

40:26The distinction basically kind of where in the chain of thought it's like static is like towards the end and dynamic is towards the beginning or? Let's say I ask you to solve a problem for me, like generate some training data. I take that, but maybe like somehow there's some noise in that. So like, you know, when you're communicating it to me, I like maybe accidentally flip a digit or like I accidentally miss a line. So like I'm not changing any of the global structure of your solution. I'm just like changing very like local things. That's static noise. Whereas dynamic noise, you make a mistake when you're generating the solution and then that impacts the rest of the answer.

41:08So like now the mistake that you've made has like, you know, messes up step two, step two messes up step three, all the way down the line. But is dynamic noise really noise in the training data it could be yeah like it could be that you know gpt2 tries to solve a problem and it starts from a pretty good place but then it totally veers off like it actually just like does something totally unreasonable that would be a form of dynamic noise and for like a practical form of static noise would be like maybe gpt2 knows or gpt4 knows like the high level steps for how to get to solution so it does roughly the right things but when you look at a particular step it makes a mistake which doesn't really matter but like technically speaking is wrong That would be a form of static noise.

41:49One of the main findings of this paper was when you have the static noise, so let's say you have some data set which has some level of static noise. Let's say that in every single training sample, you have some noise, okay? So we investigated different regimes. In some regimes, only 50 % of the training data set had noise. But in this regime, I'm telling you every single sample is infected. So that's already pretty bad if every single sample has noise, right? And now I'm also telling you, okay, consider the case where 70 % of all your digits in your data set are like wrong or flipped or something.

42:25Like you've corrected 70 % of the digits in your training data set. So now in every single sample, 70 % of numbers are just wrong. So now fine tune your model on that. How do you expect it to do? I would think that it would do okay. You think so? Okay. Well, you're right. So if you add noise to every single training sample and you like corrupt 70 % of the digits, it does fine. It gets 100 % accuracy, which to me was pretty shocking that you like you can add that much noise and like have it be fine. The situation breaks down. So if you if you add noise to 90 % of the digits in every single training sample, then you start to get in trouble and like your model just can't learn anything.

43:08but I still think it's like pretty remarkable that models, transformer models are so robust that you can inject that much noise and still get 100 % performance it reminds me of like you know stupid Facebook posts where like you've got some text but like the you know there are numbers instead of letters and things are backwards and stuff like that and like the brain figures that out I guess so yeah so like there's some structural thing which it's still able to learn and generalize uh so i think that was kind of the most uh how do i want to put it um sensational finding in my in my i am surprised though that there's there's some threshold or something after which it kind of falls apart in fact i could have convinced myself that uh you know more errors are better because it focuses the model on the algorithm and not the role of the numbers.

44:07Like, for example, the next extreme case of that would just be if you blanked out all of the individual numbers, like you're really focusing out on the algorithm and not the numbers. I agree. Like you would think that some level of noise induces more generalization. And I think that's true if the level of noise is low enough. But so we found that this is also roughly the case if you delete lines. So you can delete something like 30 % of all the lines in each solution and the model still is fine. But if you delete more than 30%, then accuracy starts to go down. Dynamic noise was much more destructive, as you can probably guess.

44:45If you have even 10 % of all the samples infected with dynamic noise, then you can't really learn anything. You're in real trouble. What I'm most optimistic for is And I think, you know, lots of other people have more or less said this, but if you take some LLM and then you combine it with some exterior system that allows it to do, you know, like more complicated planning or like, in my case, allows it to explore and generate much more diverse types of solutions. I think that's the most interesting thing. So I'm skeptical that you're going to be able to get everything you want just by like using the LLM itself.

45:22but when you combine it with these other types of systems which somehow broaden the type of solutions it's able to discover i think that's where you're really going to get like you know superhuman reasoning performance that people are looking for well alex thanks so much for taking the time to share a bit about what you're working on very interesting stuff of course it was a lot of fun

45:51Thank you.

From the publisher

Today we're joined by Alex Havrilla, a PhD student at Georgia Tech, to discuss "Teaching Large Language Models to Reason with Reinforcement Learning." Alex discusses the role of creativity and exploration in problem solving and explores the opportunities presented by applying reinforcement learning algorithms to the challenge of improving reasoning in large language models. Alex also shares his research on the effect of noise on language model training, highlighting the robustness of LLM architecture. Finally, we delve into the future of RL, and the potential of combining language models with traditional methods to achieve more robust AI reasoning.

The complete show notes for this episode can be found at twimlai.com/go/680.

More from The TWIML AI Podcast (formerly This Week in Machine Learning & Artificial Intelligence)

All 156 episodes
Teaching Large Language Models to Reason with Reinforcement Learning with Alex Havrilla - #680The TWIML AI Podcast (formerly This Week in Machine Learning & Artificial Intelligence) · 46 min
Listen in VO