SPEAKER_01
Hi everyone, my name is Max. I am VP of Research and Development at Together AI.
SPEAKER_01
And today I'm going to tell you about our research project which is called Road to 5 million sequence length, breaking memory barriers in context parallelism. So to begin, I'll first say a few words about Together AI and who we are. Together AI is an AI native cloud which provides services and infrastructure for AI developers and builders at all stages of development, starting from creating a model where you just might need a GPU cluster with heavily optimized computes and highly reliable systems all the way through model shaping where you can take existing models and customize them for your tasks in terms of performance, in terms of speed, in terms of quality through services such as fine tuning or reinforcement learning. Also we are an inference provider. So if you have an app which is reliant on open source model inference, you can work with us and we'll provide you with the fastest way to launch and use AI models with more than 200 models in our portfolio, options for deployment which include serverless and dedicated inference, and a ton of advanced optimizations which I will not be able to speak about today.
SPEAKER_01
The purpose of this talk is focused on model training, customization and fine tuning in particular. And I'll start by asking a question. In the last few months or at least a year or so, we're seeing a lot of interest in the community both on the system side and on the research side in training long context models. The primary reasons for that are twofold, I would say. First of all, with the explosion in popularity of agents, you can see a lot of different applications where you might want to put as many tokens as you want in your context, and you want the model to leverage that context effectively. Second, with the development of applications such as video generation, you might often need to keep track of multiple different frames or even multiple frames per second, which can occupy quite a few tokens in your context pretty quickly. And you also need models that have good sense of temporal consistency, which means that they are able to see what was happening a few seconds or ideally a few minutes ago.
SPEAKER_01
To do that all effectively, you need to make sure that the models are able to process that context and work with it correctly at the training time. But even if you're not at the scales of millions of tokens in the context length, it's still quite important to understand where the memory goes because who knows, maybe you might be able to reinvest it in some other ways and speed up your training overall. So the problem here is that if you are taking a standard transformer based language model and trying to extend this context, you can run into two bottlenecks. Bottleneck number one is that you are faced with quadratic computation because long story short for transformer based models, you have pairwise interactions across all the elements in a sequence. The second problem is more insidious, one might say. As you continue scaling your context, your memory keeps growing linearly, which is not as bad, but still pretty difficult to deal with, unless you apply a range of specific techniques. And this is an example from Hagenfaz's blog post on model training, which shows that the sequence length growth can affect your memory limits pretty considerably.
SPEAKER_01
Here's the slide. And our goal of that project was to see how far exactly we are able to get to do that with a range of existing techniques that are pretty well known to some in the community, as well as some further optimizations that we wanted to leverage to push this a bit further ahead. So let's say you're taking a model which is a standard LLAMA 3B architecture, you're trying to fit 3 million training tokens into your context, and you're taking all of this on an 8x H100 GPU node. The first stage you'll see is that even with just model parameters, you're not able to fit it into the GPU. You run out of memory just by trying to place the model. Of course, the next stage is to apply fully sharded data parallelism, where all the parameters are chunked across the 8 GPUs that you have, which is great, but still doesn't solve the problem. You see that the memory usage for the model drops quite significantly, but you still are running out of memory because of all the attention activations. The next point that we've leveraged, and I encourage you to use as well, the next step is taking advantage of context parallelism. In particular, there is a pretty well known technique called DeepSpeed Ulysses, first introduced by Microsoft. The idea is that instead of computing all of your multi-head attention on every GPU separately for the whole sequence, you can do something more clever. In particular, you can try to compute the attention for different heads at different points in time, or on different GPUs, through communicating these activations as they are required.
SPEAKER_01
In such a way that one GPU is only responsible for one attention head here, but it's still computing the attention over the whole sequence. That technique is quite effective at addressing the problem, and it also allows it to utilize the best possible attention implementation, like flash attention one, two, three, four, to optimize that part of the computation. And then you aggregate the results as you would have previously. So if you apply Ulysses context parallelism, the utilization drops quite significantly, approximately 8x here, as it should. But we are still quite far from our goal of being able to fit that onto just a single H100 node.
SPEAKER_01
So what happens next is that we can try to recompute the activations as they are needed to us at the backward pass. That technique is known as activation checkpointing, and it's available in pretty much all of the deep learning frameworks these days that you could use. You just need to enable it in a correct way that does not impose too much of a computational burden on you. With that, with activation checkpointing, you can drop the activation usage by a further factor of 8, but still something else needs to be done. The next optimization is also connected to the storage of activations. You can try to store some of the inputs to each transformer block, not on the GPU, but instead upload them to CPU when they are not required. This is not very impactful for the performance, because you can upload it and prefetch when you are trying to backpropagate to the corresponding layer. This optimization, to the best of our knowledge, was first implemented by Onslaught, and it allows you to drastically expand the context window.
SPEAKER_01
The next point is that you are getting with uploading next to 37 gigabytes of data, but then comes the other part of out of memory usage. What happens next is that you can essentially tile all of the computations across the sequence length in case they are element-wise. So all the loss computations, all the MLPs, they can be chunked to avoid creating these huge buffers that would be 3 million along one of the dimensions. That's our sequence length training, and even with these optimizations, you are finally getting to a point where 3 million is possible.
SPEAKER_01
But what if you wanted to go further? And here we actually need to do something else, which is the primary optimization we've done in our work, dubbed the untitled Ulysses. And you could describe it as a further, deeper analysis and expansion of this context parallelism technique. So what we found was that even trying to compute one set of heads at a time is already enough to saturate the computational capacity of the GPU within one iteration. Which means that if you have multiple different heads scheduled to be executed on one GPU, you can divide it in chunks and then essentially iterate through these chunks over time.
SPEAKER_01
So we have one group of heads which are being recomputed, then you compute attention over them, you store the partial result, then you follow up with the next stage, which can reuse all the buffers you've allocated at the previous stage. So the advantage here is that instead of allocating this huge buffer as you would have before, like here in this slide, you allocate a buffer which is smaller, but you reuse it across two or more different iterations. And that allows you to save on the activation memory for your training without any significant impact to the throughput at smaller scales.
SPEAKER_01
So here you can see the results that we've measured across different context parallelism techniques. As you can see both at the 8 billion scale and at the 32 billion scale, we are matching quite closely the most memory optimized implementations of transformer training while being able to scale even further, like 5 million tokens. And sometimes even being more performant at shorter sequences.
SPEAKER_01
The relation between the chunk size or the number of heads you compute at the same time and the throughput is quite straightforward. So if your chunk is larger, your memory utilization is higher, but at the same time you can run the whole model a bit faster. So by stacking all of these techniques together and applying UPipe on top, you can, for example, free up a bit of additional memory in your training if you need it and reinvest it somewhere else, for example, among the stages. Or you could say that we're interested in training across 3 by 5 million context lengths. And UPipe is the technique that will save you by contrast to everything else.
SPEAKER_01
So as a takeaway, I think one of the things that could be quite insightful here is that training models with large context lengths is a very interesting and challenging goal, but the bottlenecks might appear where you least expect. So tooling like the PyTorch profiler, which we elaborate on in our paper, or other techniques can help you a lot. And also check out our paper for more results. All of that is public at the moment. And we have an upcoming thread which will illustrate the method in more depth. Thank you very much for listening. And now we are ready for questions. Thank you. Do you guys have any questions?
SPEAKER_01
I'm just curious about the QKV, there was quantization parameters that I correctly understood?
SPEAKER_01
Not exactly. It was just the query key and value matrices of the transformer layer. So you multiply them here in the attention part, which creates most of the complexity. Because all of the queries have to result in all of these pairwise active interactions with keys. And the problem is that if you have a sequence which is three million in length, it means that technically in the standard most vanilla way, you would have just allocated that whole big tensor, which has three million, which is three million in size along one of the axes. And that's pretty significant as you could imagine, which means that you have to resort to not just one technique, which is U-pipe, but a range of other approaches to somehow help you leverage, actually, execute these computations without running out of your memory.
SPEAKER_01
So, yeah, that's the key idea and the key challenge of working with transformers at this scale. Cool. In that case, thank you very much for questions and for listening. Thank you. Thank you. To do that all effectively, you need to make sure that the models are able to process that context and work with it correctly at the training time. But even if you're not at the scales of like millions of tokens in the context length, it's still quite important to understand where the memory goes because who knows, maybe you might be able to reinvest it in some other ways and speed up your training overall.
SPEAKER_01
So the problem here is that if you are taking a standard transformer based language model and trying to extend this context, you can run into two bottlenecks. Bottom like number one is that you are faced with quadratic computation because long story short for transformer based models, you have pairwise interactions across all the elements in a sequence. The second problem is more insidious, one might say. As you continue scaling your context, your memory keeps growing linearly, which is not as bad, but still pretty difficult to deal with, unless you apply a range of specific techniques. And this is an example from Hagenfaz's blog post on model training,
SPEAKER_01
which shows that the sequence length growth can affect your memory limits pretty considerably. Here's the slide. And our goal of that project was to see how far exactly we are able to get to do that with a range of existing techniques that are pretty well known to some in the community, as well as some further optimizations that we wanted to leverage to push this a bit further ahead. So let's say you're taking a model which is a standard LAMO3B architecture, you're trying to fit 3 million training tokens into your context, and you're taking all of this on an 8x8100GPU node. The first stage you'll see is that even with just model parameters,
SPEAKER_01
you're not able to fit it into the GPU. You run out of memory just by trying to place the model. Of course, the next stage is to apply fully sharded data parallelism, where all the parameters are like, basically chunked across the 8GPUs that you have, which is great, but still doesn't solve the problem. You see that the memory usage for the model drops quite significantly, but you still are running out of memory because of all the attention activations. The next point that we've leveraged, and I encourage you to use as well, the next step is taking advantage of context parallelism. In particular, there is a pretty well known technique called DeepSpeed Ulysses,
SPEAKER_01
first introduced by Microsoft. The idea is that instead of computing all of your multi-head attention, like on every GPU, separately for the whole sequence, you can do something more clever. In particular, you can try to compute the attention for different heads at different points in time, or like on different GPUs, through communicating these activations as they are required. In such a way that one GPU is only responsible for one attention head here, but it's still computing the attention over the whole sequence. That technique is quite effective at addressing the problem, and it also allows it to utilize the best possible attention implementation,
SPEAKER_01
like flash attention one, two, three, four, to optimize that part of the computation. And then you aggregate the results as you would have previously. So if you apply Ulysses context parallelism, the utilization drops quite significantly, like approximately 8x here, as it should. But we are still quite far from our goal of being able to fit that onto just a single H100 node. So what happens next is that we can try to recompute the activations as they are needed to us at the backward pass. That technique is known as activation checkpointing, and it's available in pretty much all of the deep learning frameworks these days
SPEAKER_01
that you could use. You just need to enable it in a correct way that does not impose too much of a computational burden on you. With that, with activation checkpointing, you can drop the activation usage by a further factor of 8, but still something else needs to be done. The next optimization is also connected to the storage of activations. You can try to store some of the inputs to each transformer block, not on the GPU, but instead upload them to CPU when they are not required. This is not very impactful for the performance, because you can upload it and prefetch when you are trying to backpropagate to the corresponding layer.
SPEAKER_01
This optimization, to the best of our knowledge, was first implemented by Onslaught, and it allows you to drastically expand the context window. The next point is that you are getting with uploading next to 37 gigabytes of data, but then comes the other part of out of memory usage. What happens next is that you can essentially tile all of the computations across the sequence length in case they are element-wise. So all the loss computations, all the MLPs, they can be chunked to avoid creating these huge buffers that would be 3 million along one of the dimensions.
SPEAKER_01
That's our sequence length training, and even with these optimizations, you are finally getting to a point where 3 million is possible. But what if you wanted to go further? And here we actually need to do something else, which is the primary optimization we've done in our work, dubbed the untitled Ulysses. And you could describe it as a further, deeper analysis and expansion of this context parallelism technique. So what we found was that even trying to compute one set of heads at a time is already enough to saturate the computational capacity of the GPU within one iteration.
SPEAKER_01
Which means that if you have multiple different heads scheduled to be executed on one GPU, you can divide it in chunks and then essentially iterate through these chunks over time. So we have one group of heads which are being recomputed, then you compute attention over them, you store the partial result, then you follow up with the next stage, which can reuse all the buffers you've allocated at the previous stage. So the advantage here is that instead of allocating this huge buffer as you would have before, like here in this slide, you allocate a buffer which is smaller, but you reuse it across like two or more different iterations.
SPEAKER_01
And that allows you to save on the activation memory for your training without any significant impact to the throughput at smaller scales. So here you can see the results that we've measured across different context parallelism techniques. As you can see both at the 8 billion scale and at the 32 billion scale, we are matching quite closely the most memory optimized implementations of transformer training while being able to scale even further, like 5 million tokens. And sometimes even being more performant at shorter consequences. The relation between the chunk size or the number of heads you compute at the same time and the throughput is quite straightforward.
SPEAKER_01
So if your chunk is larger, your memory utilization is higher, but at the same time you can run the whole model a bit faster. So by stacking all of these techniques together and applying QPipe on top, you can, for example, free up a bit of additional memory in your training if you need it and reinvest it somewhere else, for example, for example, among the stages. Or you could say that we're interested in training across not 3 by 5 million context lengths. And then UPipe is the technique that will save you by contrast to everything else.
SPEAKER_01
So as a takeaway, I think one of the things that could be quite insightful here is that training models with large context lengths is a very, interesting and challenging goal, but the bottlenecks might appear where you least expect. So tooling like the PyTorch profiler, which we elaborate on a ton in our paper, or other techniques can help you a lot. And also check out our paper for more results. All of that is public at the moment. And we have an upcoming thread which will illustrate the method in more depth. Thank you very much for listening. And now we are ready for questions.
SPEAKER_01
Thank you.
SPEAKER_01
Do you guys have any questions? Because we're together we have employees.
SPEAKER_01
So I'm just curious about the qKV, there was quantization parameters that correctly understood? Not exactly. It was just the query key and value matrices of the transformer layer. So you multiply them here in the attention part, which creates most of the complexity. Because all of the queries have to result in all of these pairwise active interactions with keys. And the problem is that if you have a sequence which is like three million in length, it means that technically like in the standard most vanilla way, you would have just like allocated that whole big tensor, which has three million in, which is three million in size along one of the axis.
SPEAKER_01
And like that's pretty significant as you could imagine, which means that you have to result to like not just one technique, which is U-pipe, but a range of other approaches to somehow help you like leverage, like actually, execute these computations without running out of your memory. So, yeah, that's the key idea and the key challenge of working with transformers at this scale.
SPEAKER_01
Cool. In that case, thank you very much for questions and for listening.
SPEAKER_01
Thank you. Thank you.