AI Engineer

Road to 5 Million Tokens: Breaking Barriers in Long Context Training — Max Ryabinin, Together AI

2544 summary words 11 min summary Watch video

Start with the signal

11 min read

Summary

At-a-Glance

  • Verdict: Watch fully
  • Core thesis: Together AI achieved 5 million token context length training on 8x H100 GPUs by stacking existing memory optimization techniques and introducing 'UPipe' (chunked attention head computation), overcoming quadratic compute and linear memory bottlenecks in transformer training.
  • Why it matters: Demonstrates practical path to training long-context models for agents and video generation at scale; the optimization stack is reproducible and addresses memory bottlenecks that appear unexpectedly in production training runs.
  • Best use: Study as technical reference for scaling context windows in transformer training; understand the incremental optimization stack (FSDP → Ulysses → activation checkpointing → CPU offloading → tiling → UPipe) and where memory goes in long-context scenarios.

Executive Summary

Max Ryabinin, VP R&D at Together AI, presents research on training models with up to 5 million token context lengths on a single 8x H100 node. The work is motivated by two key use cases: AI agents that need massive context windows, and video generation models requiring temporal consistency across hundreds of frames per second. The core problem is that standard transformers face quadratic computation complexity and linear memory growth as context scales, making even 3 million tokens impossible to fit on GPUs without specific techniques.

The solution is a layered optimization stack. Starting with a LLAMA 3B architecture on 8x H100, the team applied: (1) Fully Sharded Data Parallelism (FSDP) to chunk parameters across GPUs, (2) DeepSpeed Ulysses context parallelism to compute different attention heads on different GPUs while maintaining full-sequence attention, (3) activation checkpointing to recompute activations during backward pass rather than storing them, (4) CPU offloading of transformer block inputs when not needed (pioneered by Onslaught), and (5) tiling of element-wise operations like MLPs and loss computation to avoid creating 3-million-element buffers.

Even with these five techniques, they hit limits. The key innovation is 'UPipe' (called 'untitled Ulysses' in the talk), which further chunks attention head computation within a single GPU. Instead of computing all assigned heads at once and allocating large buffers, UPipe iterates through head chunks sequentially, reusing smaller buffers across iterations. This saves activation memory without significantly impacting throughput at smaller scales. The technique is a deeper expansion of Ulysses parallelism, recognizing that even one set of heads can saturate GPU compute capacity.

Results show that at both 8B and 32B parameter scales, the optimization stack matches or exceeds memory-optimized transformer implementations in throughput while enabling 5 million token contexts. The chunk size vs. throughput tradeoff is straightforward: larger chunks use more memory but run faster. Ryabinin emphasizes that bottlenecks appear where least expected, and profiling tools like PyTorch Profiler are critical. The work is public, with a detailed paper and upcoming implementation thread.

Key Takeaways

  • Claim: Long-context training (millions of tokens) is driven by two use cases: AI agents needing massive context and video generation models requiring temporal consistency across many frames per second. | Evidence: Ryabinin states agents want 'as many tokens as you want in your context' and video generation 'might 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.' | Caveat: No specific numbers given for how many tokens typical agent or video tasks require; 'millions' is asserted but not justified with concrete examples. | Implication: If Ken is building agents or video models, understanding long-context training techniques becomes strategically important; if not, the techniques still apply to freeing memory for other training optimizations. | Timestamp: timestamp unavailable
  • Claim: Standard transformers face two bottlenecks at scale: quadratic computation from pairwise token interactions and linear memory growth that becomes unmanageable without intervention. | Evidence: Ryabinin explains 'for transformer based models, you have pairwise interactions across all the elements in a sequence' (quadratic) and 'your memory keeps growing linearly, which is not as bad, but still pretty difficult to deal with' (linear). | Caveat: He doesn't provide runtime or memory growth numbers; the linear memory problem is described as 'insidious' but not quantified. | Implication: Even teams not targeting millions of tokens should audit where memory goes in training—freed memory can be 'reinvested' to speed up training or increase batch size. | Timestamp: timestamp unavailable
  • Claim: DeepSpeed Ulysses context parallelism distributes attention head computation across GPUs, achieving ~8x memory reduction by computing different heads on different GPUs while each GPU still processes the full sequence. | Evidence: Ryabinin describes Ulysses as 'one GPU is only responsible for one attention head here, but it's still computing the attention over the whole sequence' and observes 'utilization drops quite significantly, approximately 8x here, as it should.' | Caveat: No throughput impact numbers provided for Ulysses alone; compatibility with Flash Attention 1/2/3/4 is mentioned but not benchmarked. | Implication: Ulysses is a proven technique for context parallelism; Ken's teams can adopt it immediately when scaling context beyond single-GPU limits. | Timestamp: timestamp unavailable
  • Claim: Activation checkpointing and CPU offloading of transformer block inputs are necessary intermediate steps, but neither alone closes the gap to 3 million tokens; tiling element-wise ops (MLPs, loss) is also required. | Evidence: Ryabinin shows a progression: activation checkpointing drops usage 'by a further factor of 8,' CPU offloading (first by Onslaught) 'allows you to drastically expand the context window,' and tiling 'avoids creating these huge buffers that would be 3 million along one of the dimensions.' | Caveat: CPU offloading is 'not very impactful for performance' because of prefetching, but no latency or throughput numbers are given; the 37 GB figure is mentioned but not contextualized. | Implication: CPU offloading is a 'free' optimization if implemented correctly with prefetching; tiling is a standard systems trick that Together applied to long-context training. | Timestamp: timestamp unavailable
  • Claim: UPipe (Together's novel technique) chunks attention head computation within a single GPU, reusing smaller buffers across iterations instead of allocating one large buffer, enabling 5 million token contexts with minimal throughput impact. | Evidence: Ryabinin explains '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,' so UPipe 'divides it in chunks and then essentially iterates through these chunks over time,' reusing buffers. Results show matching or exceeding memory-optimized implementations at 8B and 32B scales. | Caveat: Throughput impact is described as 'not significant at smaller scales' but no benchmark comparison vs. vanilla Ulysses is provided; tradeoff is 'if your chunk is larger, your memory utilization is higher, but at the same time you can run the whole model a bit faster.' | Implication: UPipe is the key unlock for scaling beyond 3 million tokens; Ken's teams should watch for the 'upcoming thread' with implementation details if targeting extreme context lengths. | Timestamp: timestamp unavailable

Detailed Brief

Problem: Why Long Context and Why Memory Matters

  • Claims: Agents and video generation models need millions of tokens in context to function effectively; Even if not targeting millions of tokens, understanding memory usage can speed up training by freeing resources for reinvestment; Transformers have quadratic computation (pairwise interactions) and linear memory growth as context scales
  • Evidence: Agents want 'as many tokens as you want in your context' and video models track 'multiple frames per second' requiring many tokens; Hagenfaz blog post cited as showing 'sequence length growth can affect your memory limits pretty considerably'; No specific numbers or graphs provided beyond the conceptual explanation
  • Caveats: Use cases are asserted but not quantified (how many tokens do agents actually need?); Quadratic computation problem is mentioned but not addressed in this work (focus is memory); Linear memory growth is 'not as bad' as quadratic compute but still 'pretty difficult to deal with'
  • Implications: Ken's agent and video teams should prioritize long-context training infrastructure; Memory profiling is valuable even for standard-context training to find reinvestment opportunities; The research is defensive: bottlenecks 'might appear where you least expect'

Optimization Stack: From FSDP to UPipe

  • Claims: LLAMA 3B with 3 million tokens on 8x H100 fails even with just model parameters loaded; FSDP (Fully Sharded Data Parallelism) chunks parameters across GPUs, dropping model memory significantly but still running out due to attention activations; DeepSpeed Ulysses context parallelism distributes attention heads across GPUs, achieving ~8x memory reduction; Activation checkpointing recomputes activations during backward pass, saving another factor of 8; CPU offloading of transformer block inputs (Onslaught technique) is 'not very impactful for performance' due to prefetching; Tiling element-wise ops (MLPs, loss) avoids 3-million-element buffers; UPipe chunks attention head computation within a single GPU, reusing buffers across iterations
  • Evidence: Ryabinin walks through each technique as a stage in a progressive memory reduction diagram; Ulysses 'utilization drops quite significantly, approximately 8x here, as it should'; Activation checkpointing 'can drop the activation usage by a further factor of 8'; CPU offloading 'to the best of our knowledge, was first implemented by Onslaught'; UPipe 'allocates a buffer which is smaller, but you reuse it across two or more different iterations'
  • Caveats: No throughput numbers provided for individual techniques except UPipe results at the end; Activation checkpointing must be 'enabled in a correct way that does not impose too much of a computational burden'; CPU offloading is only 'not very impactful' if prefetching is done correctly; UPipe throughput impact is 'not significant at smaller scales' but no definition of 'smaller' given
  • Implications: Ken's teams can adopt this stack incrementally: FSDP and Ulysses are well-known, CPU offloading and tiling are straightforward, UPipe is the novel addition; Profiling (PyTorch Profiler recommended) is essential to identify which technique applies to which bottleneck; The stack is reproducible: paper is public, implementation thread is upcoming

UPipe: The Key Innovation

  • Claims: Even computing one set of attention heads at a time saturates GPU compute capacity; UPipe divides head computation into chunks and iterates through them, reusing smaller buffers; UPipe matches or exceeds memory-optimized transformer implementations at 8B and 32B scales while enabling 5 million token contexts; Chunk size vs. throughput tradeoff is straightforward: larger chunks = higher memory, faster model
  • Evidence: Ryabinin explains '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'; Results slide shows performance at 8B and 32B scales 'matching quite closely the most memory optimized implementations of transformer training while being able to scale even further, like 5 million tokens'; Sometimes 'even being more performant at shorter sequences'; Tradeoff: 'if your chunk is larger, your memory utilization is higher, but at the same time you can run the whole model a bit faster'
  • Caveats: No head-to-head benchmark vs. vanilla Ulysses provided; Throughput impact described qualitatively ('not significant at smaller scales') but not quantified; Chunk size tuning is required; no guidance on optimal chunk size selection; The term 'UPipe' is used inconsistently ('untitled Ulysses' earlier, 'UPipe' later)
  • Implications: UPipe is the unlock for 5 million tokens; teams targeting extreme context should prioritize implementing it; The technique is a systems-level optimization (buffer reuse) rather than an algorithmic change; Ken's teams should wait for the 'upcoming thread' with implementation details before attempting reproduction

Results and Takeaways

  • Claims: Together achieved 5 million token context training on 8x H100; Performance at 8B and 32B scales matches or exceeds memory-optimized baselines; Bottlenecks appear unexpectedly; profiling tools are critical; Work is public with paper and upcoming implementation thread
  • Evidence: Ryabinin presents results slides showing 8B and 32B performance; Emphasizes 'bottlenecks might appear where you least expect' and recommends 'tooling like the PyTorch profiler, which we elaborate on a ton in our paper'; Paper is public, thread is 'upcoming'
  • Caveats: No actual throughput numbers, training time, or cost analysis provided; No comparison to alternative approaches like sparse attention or linear attention mechanisms; No discussion of quality/performance of models trained at 5 million tokens; Implementation thread is not yet available
  • Implications: Ken's teams can adopt the stack today for context scaling; Together has validated it at production scale; Profiling is a first-order activity, not an afterthought; The work is reproducible but requires systems-level engineering; not plug-and-play; Together AI is positioning as infrastructure provider for long-context training and inference

Notable Concepts & Terms

  • DeepSpeed Ulysses: Microsoft technique for context parallelism that distributes attention head computation across GPUs; each GPU computes one head but processes the full sequence, achieving ~8x memory reduction.
  • UPipe (also 'untitled Ulysses'): Together AI's novel technique that further chunks attention head computation within a single GPU, reusing smaller buffers across iterations to enable 5 million token contexts without significant throughput impact.
  • Activation checkpointing (also 'gradient checkpointing'): Technique to recompute activations during backward pass rather than storing them, trading compute for memory; achieves ~8x memory reduction but requires correct implementation to avoid excessive compute overhead.
  • CPU offloading (Onslaught technique): Storing transformer block inputs on CPU when not needed, prefetching them during backpropagation; 'not very impactful for performance' if done correctly, enabling drastic context expansion.
  • Tiling (element-wise operation chunking): Dividing element-wise operations (MLPs, loss computation) into chunks to avoid allocating buffers with one dimension equal to sequence length; standard systems technique applied to long-context training.
  • FSDP (Fully Sharded Data Parallelism): PyTorch technique to chunk model parameters across GPUs; necessary first step for long-context training but insufficient alone due to attention activation memory.
  • Flash Attention: Family of optimized attention implementations (versions 1/2/3/4) that Ulysses and UPipe are compatible with; enables fast attention computation while parallelizing across GPUs.

Operator Notes / Why Ken Should Care

  • If Ken's teams are building agents or video models, this optimization stack is immediately relevant for scaling context windows; Together AI has validated it at production scale (5M tokens on 8x H100).
  • Even if not targeting extreme context, the memory optimization techniques (checkpointing, CPU offloading, tiling) can free resources for faster training or larger batch sizes in standard workloads.
  • UPipe is a systems-level innovation (buffer reuse) rather than algorithmic; it requires engineering effort to implement correctly but is not a research gamble.
  • Together AI is positioning as infrastructure provider for long-context training and inference; if Ken's portfolio companies need this capability, Together is a credible vendor.
  • The work is reproducible (public paper, upcoming thread) but not plug-and-play; requires PyTorch profiling and tuning (chunk size selection, prefetching logic) to achieve reported performance.
  • No cost or training time analysis provided; unclear if 5M token training is economically viable or just technically feasible.
  • No discussion of model quality at 5M tokens; the work is purely systems/infrastructure focused, not about whether models actually benefit from such long contexts.

Watch Map

  • timestamp unavailable: Timestamps not provided in transcript; video is 15:50 total. Structure: 1) Together AI intro, 2) Problem statement (quadratic compute, linear memory), 3) Optimization stack walkthrough (FSDP → Ulysses → checkpointing → CPU offload → tiling → UPipe), 4) UPipe deep dive, 5) Results and takeaways, 6) Q&A (one question on QKV).

Source/Metadata

  • Title: Road to 5 Million Tokens: Breaking Barriers in Long Context Training — Max Ryabinin, Together AI
  • Transcript words: 3565
  • Duration seconds: 950
  • Timestamp note: Timestamps unavailable; transcript structure suggests ~2 minutes intro, ~8 minutes optimization walkthrough, ~3 minutes UPipe/results, ~2 minutes Q&A.
Full transcript 1993 words · 16 min read
0:14

SPEAKER_01

Hi everyone, my name is Max. I am VP of Research and Development at Together AI.

0:19

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.

0:25

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.

0:32

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.

0:38

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.

0:46

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.

0:52

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.

1:01

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.

1:06

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.

1:14

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.

1:24

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.

1:33

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.

1:48

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?

2:13

SPEAKER_01

I'm just curious about the QKV, there was quantization parameters that I correctly understood?

2:22

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.

2:29

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.

3:42

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,

4:36

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,

5:36

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,

6:23

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,

7:14

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

8:06

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.

9:05

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.

9:51

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.

10:40

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.

11:27

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.

12:19

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.

13:00

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.

13:51

SPEAKER_01

Thank you.

13:57

SPEAKER_01

Do you guys have any questions? Because we're together we have employees.

14:06

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.

14:58

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.

15:30

SPEAKER_01

Cool. In that case, thank you very much for questions and for listening.

15:45

SPEAKER_01

Thank you. Thank you.

Reading tools

Type to find a passage

Appearance
Ask this transcript

Add a note