GPUs
Goals, and why hardware decides the ceiling
00:04 So hopefully everyone's having a good time with assignment one. It's due tonight. Let us know if you need an extension. Assignment two is coming out soon. We're putting on the finishing touches onto some of the Triton stuff. Hopefully, you'll enjoy it. You'll get to implement FlashAttention-2 or parts of FlashAttention-2, which I think will be nice. So today we're going to talk about GPUs. GPUs are the thing that makes our language models go. So they're pretty critical to get right. And if you haven't really studied the hardware that makes your models run, they can seem pretty mysterious. So my goal today is to try to make CUDA and GPUs less magic. And one of the things that I want to demystify-- you don't have to understand the plot. There's a lot on the slide I know-- why do GPUs get slow. And they get slow in very mysterious ways. I will try to talk through this plot towards the end of lecture. As you increase the size of your matrix multiplies, you might expect either gets slower or faster or whatever. You get these very unpredictable looking wave like patterns and you're like, why is my GPU fast at certain multiples of certain numbers and slow at others. It's very mysterious. We'll try to understand that. The other thing is we would like to understand how to make fast algorithms. I think almost all of you have heard of FlashAttention. It's the thing that makes much longer contexts possible by very cleverly computing the attention operation inside a transformer. And so maybe you would like to come up with new algorithms or new implementations
01:47 like FlashAttention. Like what primitives and what components do we need to understand in order to be able to do that. So those are the two learning goals of today. The first one is by the end of the lecture, you should feel comfortable with GPUs. You should understand how they work. And the second one is you should feel comfortable accelerating certain parts of your algorithms. You make a new architecture, you should hopefully feel like you can try to accelerate that with CUDA. And because hardware is not necessarily the domain in which I work. There's special resources that I have to give a lot of credit to, especially Horace He's blog, where he's got a lot of fun GPU facts that you can learn about. For example, why are matrix multiplies that are filled with zeros faster than ones that are not filled with zeros, you can learn by going to his blog. There's also other resources that I've drawn from, like the CUDA MODE group and the nice TPU book from Google. If this topic interests you, I'd encourage you to go and look at those resources to learn more, because this is in some ways like a shallow but hopefully, complete coverage of the hardware. So today we're only going to focus on non-parallel parts of the hardware stack. So we're going to study the GPU, like a single accelerator in depth how they work in some important parts. I'm also going to talk very, very briefly about TPUs, because in some ways, they're very similar conceptually to a GPU. And so my discussion here is going to carry over. And then once we understand the hardware and execution model of the GPU, then we're going to try to understand what makes GPUs go fast on certain workloads, what makes them slow. We're going to understand the performance. And in the last part, this is kind of going to be almost like a hands-on piece. I'm going to try to walk through FlashAttention.
03:25 I'm going to take all the lessons that we've learned and try to walk you through FlashAttention saying, see, here's how it all comes together. So that's the last part of today's lecture. So many of you have taken an NLP course. And these days in an NLP course, I think you teach some amount of scaling laws. And so you've probably seen this, right? And so this is just setting the context. We know that having more compute is helpful for training large language models. This is a pre-training scaling chart. But you could replace this with an inference scaling chart if you would like. It's generally agreed upon that the more compute you have, the more processing you can do on your data. You can ingest more data. You can train larger models. All of those lead to improved performance. So you might think of course, deep learning is really important. But what's really driven performance is faster hardware, better utilization, improved parallelization. So that's setting the stage of why hardware is important to understand. And of course, once you think about compute scaling you ask, OK, how do we get compute scaling? How do we get our models to train faster? So in the early days, of semiconductor scaling, if you were thinking about, OK, CPUs, how do they get faster, they would scale under something called Dennard scaling. With Moore's law, you would double the amount of transistors on a chip every year. And if you have this doubling, what you end up is Dennard scaling where smaller and smaller transistors can be driven at faster and faster clock speeds with lower and lower power, which in turn give you more performance. And then in the 1980s to 2000, this tapped out. You can see in this chart here by Hennessy and Patterson that single thread performance, that's the blue dots here,
05:09 that basically started to taper out. Of course, the number of transistors didn't really start falling off. You did have chips with higher and higher transistor densities. But that wasn't helpful. It wasn't giving you higher throughput on single threads. And so, this means that we can't just do computation faster in absolute terms. What we have to make up for it with is parallel scaling. So the story of scaling for deep learning and neural networks is going from single-thread scaling, which is just doing your computation faster in absolute terms, to parallel scaling where you have a lot of workloads that are all computed at once. And this is one of my favorite compute scaling charts by Bill Dally and his keynote where he's showing the super exponential increase in the number of integer operations per second, going from the earliest K20s to the H100. And it's this really remarkable exponential or super exponential curve. And so we have to really understand how to take advantage of this curve in order to really get the most out of our language model.
CPU vs GPU: latency-optimized vs throughput-optimized
06:18 So that's kind of going to be our goal. And so I've already hinted at this kind of important difference, right? CPU is something that I think everyone is familiar with once you start doing programming. It's this execution model. If you have a program, it goes through. And in a single thread, it executes step by step what's happening. And in order to support that kind of an execution model, what do you need? Well, you need big control units. You just need to generally run these things very quickly because you have a lot of branching and you have a lot of conditional control logic. So a CPU-- this is an abstracted diagram-- is going to dedicate a lot of its chip towards large control branch prediction. And it's going to run these very quickly because it doesn't have that many threads. There are CPUs with lots and lots of cores now. But compared to a GPU it's almost nothing. And so in contrast, the GPU has really tons and tons of compute units ALUs. So those are the little green boxes. And there's much smaller amounts of the chip dedicated to control. So there's a little bit of control logic orchestrating tons and tons of compute units operating in parallel. And I think mentally-- so this is the picture of what is being emphasized in a CPU versus a GPU. But if you look at what's the design goals are, they designed for very different goals. So you can think about CPUs as optimizing for latency. I want to finish my tasks as quickly as possible. So if I have tasks T1 through T4 here on the right side, in a CPU, I'm going to try to finish each task as quickly as possible. And so if you want any one of these tasks to be finished quickly, T1 is going to complete really quickly.
07:59 In GPU, you're optimizing for high throughput. I don't care about latency, I just want all of my tasks that I have in aggregate, to complete as quickly as possible. And to support that, maybe you have lots of threads and these threads can go to sleep and wake up very quickly. And in the end, you finish all of your workload T1 through T4 before the CPU 1 does, even though individually all of these have higher latency. So they have different design principles and design goals. OK. And so a GPU has a pretty different anatomy.
Anatomy: SMs, SPs, and how far away the memory is
08:33 And I don't know if you all have ever looked at what a GPU layout diagram looks like. I'll actually show you the chip figures in a moment here. But the core idea, and this is important conceptual concepts behind a GPU, is that a GPU executes many, many SM streaming multiprocessors. And a streaming multiprocessor you can think of as an atomic unit. When you're programming in something like Triton, they're going to operate at the level of an SM. And within each SM, they're going to-- it contains many SPs, Streaming Processors, and a streaming processor is going to execute a whole bunch of threads in parallel. So one way to think about it is SM has a bunch of control logic. It can decide what to execute. It can do for example, branching. SPs are going to operate-- to take the same instruction and apply it to many different pieces of data. And so you can do tons and tons of parallel computation. Under this model, an SM is each granular unit of control. SP can do a lot of computation individually. And if you look at an A100, which is the previous generation GPU, at this point you've got 128 SMs. That's a lot more than most cores for CPUs. And each of these systems is going to have a very large number of SPs and specialized matrix multiply units inside them. And so that's kind of the compute model. Was there a question? Sorry. Yeah. [INAUDIBLE] back to the other slide before Anatomy of GPUs. So is this GPU the same as that GPU? So the question was, is this GPU the same as that GPU?
10:13 Yes. Like this is a kind of cartoon version of this. You can think of each row as being an SM. It's got its own control units, and each green block might be one of these green blocks here, like a FP32 processing unit inside of it. And each SM can operate various pieces that it owns, like the tensor cores to do computation. Cool. OK. And there's going to be two important things. You think of GPUs as computers. They compute. But actually computation is only one of the two important things we have to keep track of. Memory is arguably more important, at this point, and it will continue to be more important in terms of the performance profiles of how we run our programs on the GPU. And so to understand memory, you have to understand the physical layout of the GPU and the chip, because in some sense, when you're operating at such fast speeds, the physical proximity of the memory starts to matter quite a bit. And so I will show you the physical proximity of how things are laid out and how that relates to how you should think about memory access and performance. So the closer a piece of memory is to each SM, the faster it's going to be. So there's going to be certain very, very, very fast kinds of memory, like L1 and shared memory, and that's going to live inside of the SM. And that's going to be really fast. Things like registers, things you're reading and writing very frequently, you're going to want to put into the L1 and shared memory. L2 cache, as you can see, there's these green areas which are SMs. And then there's these blue areas. This is on the GPU chip.
11:50 These are L2 memory that's right next to the SMs. So they're not inside the SMs, but they're physically still quite close. And these are still pretty fast. They're still a factor of 10 slower, but they're still reasonably fast. And then outside of the chip itself, I think this is like a 3090 card or something like this, or maybe a PCIe A100. Oh, this is a PCIe A100. You've got your GPU here, and you've got actually DRAM living next to the chip. So it has to actually go physically outside of the chip and connect. And you can see on this chip diagram here these yellow connectors at the edges. These are HBM connectors. These are connecting to the DRAM chips that are outside of the actual GPU. And you can see the speed that it takes to access these, right? The on SM memory is much, much faster, like 20 clock cycles to access something from there, whereas it's going to take something like 200 or 300 clock cycles to access something from the L2 cache or global memory. And this factor of 10 is going to hurt you real bad. So if you have a piece of computation that requires you to access global memory, it might mean that you actually run out of work to do on your SM. You've multiplied all the matrices. You've run out. Now you just have to idle. So you utilization won't be good. And this will be a really key theme, thinking about memories in some sense. The key to thinking about how GPUs work. And in assignment two, you're going to actually be writing high performance code for a GPU. So you have to actually think about the execution model of how a GPU actually executes things.
Execution and memory model: blocks, warps, threads
13:29 And this is somewhat complicated, but not insanely so. There's three granularities of things that you need to think about. There's blocks, there's warps, and there's threads, and that's the order in which the granularity narrows down. Blocks are these big groups of threads. And each block is going to be assigned to a SM. So think about this as each SM is kind of a worker. It's its own autonomous unit. And a block is going to be assigned to an SM to process. So this is each granular unit. Now then within these blocks are a whole bunch of threads. Each thread is a piece of task that needs to be done. And when these threads execute, they're going to execute in groups. And this is a thing called a warp. So you take a block which is a collection of threads. And you're going to take threads from that block, and they're going to execute in groups of 32 consecutively numbered threads each time. And that's of called warps. And so you can see at this diagram here what's happening. You've got a bunch of blocks. Each block is assigned to a different SM. And within each block, there's going to be many different warps. And each warp is going to consist of a whole bunch of threads. And all of these threads are going to execute the same instruction on different data. And so this is kind of the execution model. Right now it's going to it seems probably mysterious. What these blocks, and warps and threads are. They will have important implications for performance in how we design things like CUDA kernels later. So hopefully you can remember this. I'll refresh your memory as we go. Hopefully, that's clear. So that was the kind of logical execution model of a GPU.
15:08 And if you understand that you understand how GPUs execute things. There's also a logical memory model of a GPU. So now I'm not showing you the physical hardware. This is just how you think about the programming of a GPU. And so there's registers. So these are really fast, storing single numbers type storage. You've got local memory, you've got shared memory, and you've got global memory. And that increases in the memory hierarchy. You get slower and slower and slower. And your code can write to global memory. It can also write the constant memory, which is not something that's used too often. And so each thread can access its own register and shared memory. But information that goes across blocks need to be written to global memory. And this is actually quite important. So now it means that whenever you write a thread that executes something, ideally it's operating on the same small amount of data. So you load that small amount of data into shared memory. All the threads are very happy accessing that shared memory. It terminates. It's done. That would be a great execution model. Instead, if you have a thread that needs to access data all over the place, that's going to have to access global memory, that's very, very slow. This theme will come back as we talk about different ways of operating on a GPU. Hopefully that's clear. That's kind of the very high level, four-slide overview of a GPU. If you have questions about how any of that works, feel free to ask me as I go on. OK, so here's a side thread. Last year I didn't cover this because I think resources on TPUs was a little thin, but the nice TPU book or internet
TPU aside, and the strengths of the SIMT model
16:45 website that I mentioned at the start of the lecture came out, and that has actually a lot of nice details. And I talked to a few Google people about the TPU. And at a high level, it's very, very similar to a GPU. And so I want to just talk for a moment about TPUs. You may never operate on a TPU, but I think it's important to understand that these alternative accelerators operate in many ways, very similarly. So here's a diagram of what a TPU looks like. There's kind of a so there's something called a tensor core. And mentally you can think about a tensor core as being similar to an SM or streaming multiprocessor. Each of these are kind of its own atomic units that can operate on data. There's a scalar unit which is basically a control unit. And it can also do CPU-like arbitrary things. You've got a vector unit that can operate on vectors. So if you've got a vector and you want to operate entrywise on it, that's a good place to do it. And then it's got a very big specialized part of the chip dedicated to just doing matrix multiply. It's called the MXU. And then it's got very fast memory for vector memory and SM. Both of these are very fast on chip or on tensor core memory. And then there's high bandwidth memory that lives outside of the chip. So hopefully you see the similarities to an SM. There's slow memory outside very fast memory inside. And there's specialized hardware to do matrix multiplication. The core structure is very much the same. The difference is I'll talk about this in the parallelism lecture next week. How the accelerators are networked together is a little bit different. And then also mention I didn't notice I didn't talk about warps, I didn't talk about any of that other stuff. Tensor cores are in some ways very simple because they're optimized to just do matrix multiplies. The tensor core, unlike the GPU, doesn't attempt
18:25 to do anything but that. And so in some ways very, very simple, much simpler in architecture, but conceptually doing the same thing. Yes. Is it called tensor because it's also in some ways optimized to operate on general tensors or this is just enough to work on [INAUDIBLE]? Yeah. So the question was is it called tensor because it can operate on arbitrary tensors. So it can operate on arbitrary tensors. I can do the indexing. The operations that MXU performs is a matrix multiply. And so it would always be like a batch matrix multiply operating on a tensor. So it's kind of both a Yes and a no answer, if that makes sense. So they operate on tensors, but the operations they always perform are matrix multiply, is not more complicated tensor operations that you can do. Cool. The reason why the GPU has been so successful is that it scales up really easily. If you want more processing power, just add more SMs. You don't have to worry about driving the clock faster and getting more heat dissipation problems. Programming-wise, CUDA is intimidating, but it's actually not as horrendous to program because of its programming model. Like, the way it works is within each SM right, you have a thread, and it executes the same instruction on a bunch of different pieces of data. That's conceptually easy to reason about. You can think through what that means. And especially it's nice if you're operating over a matrix and you're doing very simple operations. It's exactly this kind of simple model. Finally, each of these threads are very lightweight, and they can be kind of stopped and started at any time. And so if you need to wait for another thread
20:01 or if you need to evict something and start another process, all these threads are very lightweight. So this just kind of means that there's not much state associated with the threads, and they can be stopped and started, which allows GPUs to get high utilization within each SM. So GPUs, obviously graphics processing units. And for much of its life, in the early days, it was not used to do scientific computing. But people-- because it was programmable, researchers figured out how to use early NVIDIA GPUs to do
Matmuls are blessed, and the memory wall
20:37 fast matrix multiplies. This is one of the early papers on doing fast matrix multiplies with graphics hardware, and it shows how you can hack kind of things like the texture buffer. And so on to get it to do matrix multiplies. And so even without specific support for matmuls, researchers figured out how to do it. But I think now, especially in this day and age, NVIDIA and others have realized matrix multiplies are special. If you're doing deep learning, most of your workload is matrix multiplies. And so matrix multiplies are in some sense blessed operations. So this is a chart showing the number of teraFLOPs per second by different generations of NVIDIA GPUs. And the orange line is your matmul FLOPs, your performance you can get if you're doing matmuls. The blue line is your non-matmul FLOPs. And you see this big, big gap at V100s when they started putting in tensor cores that were specialized hardware to do matrix multiplies. And you see this gigantic gap in the matrix multiply performance relative to the non-matmul performance. And so if you're going to design any a neural architecture, I was saying this in the architecture part as well. You have to have most of your workload be matrix multiplies because that's the thing that's orders of magnitude faster than any other operation that you're going to be able to do on a GPU. So if you make like a non-matmul-based neural network, you're going to be in a big, big trouble. And then kind of the last thing that I want you to understand as just general facts. Matmuls is fast is one thing. But the other thing that's important to remember is kind of the relative scaling of the different components of the GPU. So this is a very nice chart that shows you
22:19 how quickly different components of the GPU or different components of the let's call it like LLM training stack are scaling. So the blue line is the connectivity from the GPU to the host. The server that it's attached to. So you can use PCIe, you can use NVLink, you can use all these fancy interconnects. They are growing, but they're growing somewhat slowly. So this chart is like normalized scaling bandwidth relative to when the first generation of interconnects. The green line, this is the global memory speed. So you go from GDDR to HBM2E, and that's much, much faster right. This is log scale. It's 100x faster. But this is still kind of slow scaling. And the gray line here, this is compute scaling. This is the number of floating point operations, if you're considering the matmul FLOPs. This is how fast the compute has been scaling. And this is astoundingly fast. It's 1 to 100,000 times faster. And so in the early days of this scaling, maybe your problems were FLOPs based. You just didn't have enough FLOPs to do your matrix multiplications. But now all the way to the right with the A100s, these are astoundingly fast GPUs. Your bottlenecks are probably going to end up being memory, because the memory is not growing as fast. And as we go into the future, this is not really going to change. DRAM is very hard to scale. You're going to keep getting this bigger and bigger gap. So if you're ever designing hardware efficient algorithms, you're going to have to think more and more about memory. And so we're going to keep a lookout on that. I'm going to keep emphasizing this. It's one of the important themes in GPUs. OK. So I've been kind of throwing lots of GPU facts at you,
24:02 especially if you haven't, seen this recently. It may be kind of new. So just to recap, GPUs are these massively parallel processing systems. They have same instructions applied across many different threads. And they have these things called SMs, which are kind of cores that there's many, many of them in the GPUs. Compute and matrix multiplies have scaled really fast, and they have scaled faster than memory. And that is an important part of the characteristics that you think about, about GPUs. But there is some fast memory. It's not like everything is slow, so there's nothing we can do. There's the memory hierarchy. So some kinds of memory are very, very fast. Other kinds of memories are slow. And so if we exploit this hierarchy maybe we can get things that are really, really fast. So that's things to remember about the GPU. And if you remember these facts, you're going to be able to think pretty cleanly about the performance components that I'm going to talk about next. Any questions before I move on to the next part? OK, cool. So now you all are GPU experts, and what we would like to do is we would like to make machine learning workloads go
The mystery plot and the roofline model
25:09 very fast on a GPU. And so I'm going to start with this chart. And one of our goals will be to understand what this chart exactly is. I think it'll be a good puzzle to get us motivated. And so here what we are doing is we are multiplying square matrices together. So the x-axis is the size of my square matrix multiplies. And the y-axis here, this is the number of operations per second that I'm doing. So you can think of this as hardware utilization on the y-axis. And so as I get bigger and bigger matrices, I'm going to get better and better hardware utilization because I have more work to do. So I don't-- that overwhelms the overhead of launching jobs and things like this. But there's all these weird things that are happening. You see 1, 2, 3 different, 4 different lines. And each of these lines are kind of wavy in a way that's looks very unpredictable. And so we would like to understand what exactly is going on with these lines. And by the end of this section, my promise is that you will understand exactly each one of these phenomena, and you'll be able to say, yeah, that plot looks totally normal. That is a natural thing for a GPU to do. So the very first part, is if you look at that plot, you will notice that it looks a little bit like this. And if you've taken a systems hardware course, you should remember this as kind of the roofline model. The roofline model basically says if we're looking at throughput or utilization, what we're going to find is there's two regimes. There's going to be a regime that is memory limited, that is on the left side of this curve
26:48 on the green over here. And then there's a part that is throughput limited on the right side. In some sense, you can think of it as, on the right side, we are fully utilizing our compute units. All the matrix multiply units are multiplying all the time. And on the diagonal here, we just have some sort of memory bottleneck. And so our ability to do computation is limited by the amount of intensity that we have, the amount of FLOPs per byte that we have. So we want to avoid being in this left side region where we're memory bound. And we would like to be on this right side where we're getting, in some sense, full utilization of all of our compute units. So that's, in some sense, the goal. And hopefully this roofline model looks something like this. We've got this diagonal part. And then we've got this flat part all the way at the top here. So that's one part of the mystery. And so this turns out to be complex. The simple way to say this is-- let's make sure that we're not accessing memory unnecessarily. We have as few memory accesses to slow global memory as possible.
Control divergence, and trick 1: low precision
27:52 But it turns out that in order to do that, we need a large array of tricks. There's a lot of different things that you could do that would mess you up, that would make you very slow. And the first one is not a memory bottleneck. I'll just mention it. It doesn't come up too often. We'll get it out of the way. And then we'll talk about the remaining five items that in some sense are really core to thinking about GPU performance. OK, so the first thing that I want to talk about is conditionals. So as I said before, GPUs, their execution model is something called SIMT, Single Instruction Multi-Thread. And so every thread in a warp is going to execute the same instruction, and it's going to do so on different data. And so what happens if I write a piece of code that looks like this. I have an if statement. And if the thread index is less than 4, do something. If the thread index is greater than or equal to 4, then do something else. I have this very simple conditional model. If I run this on the GPU, what's going to happen is that I'm going to run the A instruction on four of my threads. I will actually pause my other four threads, which are supposed to be executing the else part. And then these other four threads will come alive and they will execute x. And these my original four threads will go to sleep, and I will just alternate executing each of these instructions. Why is that? I can't execute A and x at the same time on these different threads. As I said again, every thread has to execute the same instruction. So conditional statements within a single warp can be really, really damaging because they will force you to pause any of the threads that are not doing exactly the main control flow execution. OK, so that was the only non-memory thing
29:38 that I wanted to mention, and it should be kind of obvious that you should probably not be putting conditionals into your massively parallel compute unit, but once we've gotten that out of the way, the other tricks that we need to consider are all kind of memory based. The first thing I want to mention is lower precision. And this is a big trick. This is an important trick. You should do it all the time. There's kind of a going back to this plot of Bill Dally, there's a sleight of hand here. This looks really good because the numbers are going up and up and up. But if you look at what's driving GPU progress over all these years, you actually kind of see that it's number representations. You go from FP32 to FP16 to int8 to so on. You get many orders of magnitude gains from just having lower and lower precision in your GPU operations. And let me clarify why that's so important. If you have fewer bits in all of the things that you're computing and your weights and so on, you have much fewer bits to move. So even if you're accessing these bits from global memory, they become much, much less of a concern. So let's just give a simple example. And let's just think about arithmetic intensity of a simple elementwise operation. So I'm going to do it in ReLU. So that's x equals max 0 and x. And I'm going to do that on a vector of size n. Let's say naively I'm going to do this on float 32. So how many memory accesses do I have? I have to read my x. I have to write the result of if x less than 0. And that's all in float 32. So that's kind of 8 bytes. And how many operations do I do? Well, I have to do x less than 0. So that's one comparison operation.
31:17 And I do one FLOP. So I do 8 bytes per single floating point operation. If I do this in float 16 now, well, I haven't changed the FLOPs intensity here, but I have the memory access. And so now I have 4 bytes per FLOP. In some sense I've gotten double the memory bandwidth for free, assuming that I can get away with FP16. And this is a key part of how a lot of things are designed. Part of the assignment is going to be you're going to try and play with various mixed precision or low precision training and other kinds of things. And a key part here is that not all the parts of your network and your training algorithm should be put into low precision. So let me give you an example of matrix multiplies. So in matrix multiplies that are mixed precision, what you would do is you would have your inputs be 16-bit. So these are low precision. And then you're going to do your multiplication in full 32-bit. And that's useful because the intermediate computations as you're accumulating partial sums, you would like that to be in high precision. And so you're accumulating this with FP32 accumulator. And then your tensor core will return a FP32 result, which you can downcast if you would like back into 16-bit. And so we have our inputs in 16-bit. But things like the accumulation we might want to do in 32. So there's lots of different things. There's operations that can use 16-bit storage. There's operations that might need more precision. So you want to keep it in either FP32 or FP16. You might want to have operations that need more range, like exp functions. If you don't have the dynamic range that might blow up or zero out. And so you might want to put those in bf16.
Trick 2: operator fusion
33:02 There's a lot of careful engineering that has to happen in order to make sure that these models are actually stable when they're being trained with lower precision. But if you can do it, that's really great because you've basically doubled the throughput of your bottleneck going from 32 to 16-bit, if your memory is your bottleneck. The other one, and I think this is kind of what a lot of people think of when they say I'm going to write a CUDA kernel or something, operator fusion is kind of both very intuitive and both like a fun, natural one to think about. So one memory-- or sorry, one mental model of how a GPU works and how memory works is this kind of fun diagram of a factory from Horace He. So imagine you have a factory, and your factory is your compute part. And so it takes in little box widgets and then outputs little triangle widgets. And if you grow your compute but your belt conveyor that takes memory to compute is finite bandwidth, you're not going to be able to use your second factory. You're still capped by the speed at which you can transfer things from memory to compute. And so you've got this bottleneck. Now, of course you already knew that. I've been hammering in the memory bottleneck thing, but I think one insidious way in which you can incur a ton of overhead without really realizing it is this left-hand-side computation pattern. So imagine the left side of this plot is where the memory is. The right side is your compute unit. And so to do computation I start with a square. And I move my squares from my memory to my compute. I do some operation. I turn them into triangles. Now I shift my triangles back to memory. And then I realized I need my triangles again.
34:45 So I ship them back into the compute unit. Now the triangles become circles, and so on and so forth. I send my compute back and forth and back and forth, back to memory. And you might call this a very naive approach. And if you were just doing operations naively on the GPU and just shipping the results straight back to global memory, this is what you'd end up with. And if you count the number of times a piece of data went back and forth, this is pretty terrible. You've incurred tons of memory overhead. Now, you should be able to realize that if you look at the right side, will this compute? Well, there's no dependency, so I should be able to go square to triangle to circle to rectangle and ship the rectangle back. I can just keep everything in the compute unit the whole time. And that's the right-hand-side diagram. And this is the mental model of a fused kernel. You have a bunch of operations that are going to happen on a piece of data in sequence. Instead of writing it back into storage, what I'm going to do is I'm going to do all the computation as much as I can in one place, and then only when I have to ship it back to memory. So that's this idea of kernel fusion. There's some very simple examples of how if you write some naive code, you might get of a naive set of launches. So here's an example I wrote a little let's say neural network module. Let's say I write a neural network module that takes in x and it produces sine squared x and cosine squared x. Simple code. Now if I run this, the computation graph in PyTorch is going to look something like this. And it's going to launch a whole bunch of CUDA kernels. It's going to launch take in the x, and it'll launch a CUDA kernel to compute sin x. It'll launch one to compute cosine x then sine squared of x and cosine squared of x and sine
36:27 squared x plus cosine squared of x. So there's a bunch of back and forth that has to happen in order to do this computation. It's exactly the left hand side figure that I showed you before. But if you were a little smarter and you either wrote your own CUDA kernel or you use something like torch.compile, well, you can easily realize that those five operations don't really depend on very much, they use only a little bit of memory. And so you can fuse them into a single operation that does everything on GPU on a single thread without sending things back to global memory. So really easy fusion operations like this can be done automatically by compilers. I just mentioned torch.compile. If you aren't already doing this, you should. You should consider strongly thinking about using torch.compile everywhere. We'll show you in the assignment torch.compile as well. It's pretty nice. OK, so I've gone through precision and fusion. If anyone has questions, let me know before I move on to recomputation and other kinds of tricks that we can do on the GPU.
Trick 3: recomputation
37:33 OK, good. So another thing that we can do is called recomputation. And recomputation is this idea of spending more compute to avoid having to do memory access. So remember your original back propagation lecture. This one is actually from CS221. What do we do? Well, we take our inputs at the very bottom. These are the yellow ones. And then we propagate activations upwards. Those are also the yellow values on the tree. And then we compute the Jacobians backwards. Those are the green values on the edges. And then to compute my gradients I'm going to propagate. You multiply. So the Jacobian and the activations I'm a propagate the gradients backward. Well, if you think about it, those yellow values after the forward pass have to be stored. And then they're stored. And then they have to be taken from global memory where I stored them and put them into the compute units. Mechanically, that's how it has to happen. But that might actually be a ton of memory inputs and outputs happening. Instead you might actually be able to avoid this. So let me give you an example of how computation can speed things up. Here's another silly function that I might write. I'm just going to stack three sigmoids on top of each other. You can look at the left. That's the forward graph. That should be exactly your mental model of three sigmoids on top of each other. Now the computation graph for this, I'm going to compute the sigmoids, and I'm going to store S1 and S2, which are the activations of the sigmoids. And I have my outputs. And then that's my forward pass. Now, the backward pass in this is kind of terrible.
39:12 When I do my backward graph, I need to go and take S1 and S2, and I need to take the gradients coming backwards into this out box and then push it into this backwards computation. And I'll get the gradient of x. So I need to have three memory reads one memory write in order to compute the backwards pass. And then for the forward pass I need to do one memory read of x. And I need to do three memory writes for S1, S2, and out. So hopefully that's clear. This is a decent amount of memory reads and writes. I have to do eight of them. And I have very low arithmetic intensity because I have no matrix multiplies at all. So the idea of recomputation is to say I don't want to store those activations at all. I'm not going to put them into memory. I'm just going to recompute them on the fly in my backward pass. So now in my new forward pass, I don't store S1 and S2. I take x as input, I compute my sigmoids, and I get my output. So now that's one memory read for x, one memory write for out. Now in my backward pass, I don't have activations anymore. So what I'm going to do is I'm going to get both D out, which is the backward signal coming in from above. And then x which is my input. So I'm going to take two of those, which is 2 memory reads. And then on the fly in my SM in my local memory, I'm going to compute each of these sigmoids, and I'm going to put them into the backward graph. I'm going to recompute S1, S2, and out on the fly inside my local memory. And because I do that, there's no global memory reads happening here. And then I have one memory write which is dx. So now if you compare the two I have 5/8 of the memory access for the exact same computation.
40:51 The price that we paid is that I'm going to have to recompute these three sigmoids. But if you were running idle anyway because you were memory capped, this is a great trade off. You would be very happy with this because now you've traded compute, which you have too much of for memory bandwidth, which you had too little of. So this is one great way of trading one thing you need for another thing that you have. And of course, this is different. It's the same trick as gradient checkpointing and recomputing activations for memory savings.
Trick 4: burst mode and memory coalescing
41:23 But this is being done for different reasons. This is for execution speed, not just because you're running out of memory. So it's the same technique but for different goals. And then this one I think is actually kind of a really interesting one, and not one that I knew until I started really looking into how the hardware model of a GPU and DRAM works. So the slow memory, the global memory called DRAM in a GPU, that's actually very, very slow. And in order to make it faster, there are certain optimizations that are being done at the hardware level. And one of the optimizations that's done at a hardware level for DRAM is that when you go and read a piece of memory, you don't actually get just that value back. You actually get a whole chunk of the memory back. And this is called burst mode. So let's say I went on and tried to read the very first value of this big memory block. Instead of just the memory giving me back a 0, it would actually give me back 0, 1, 2, 3. It would give me back four values at once. It would be like, here you go. I'm sure you'll need the 1, 2 and 3 too in the future. And so each address space is cut up into what's called burst sections. And then you're given the entire burst section rather than just what you looked for. And this might seem very mystifying why would the memory give you three extra bytes for free when you're just asking for one. There's a very interesting hardware reason, which is that when you're addressing into the memory, in order to send the signal out from the memory, that those bytes have to be moved to an amplifier, that's the slow step. And once you've done that, you can get many, many bytes bites for free.
43:03 And so that's why this burst section thing exists. It's kind of masking this more expensive step of actually moving where the data is stored to this amplifier. But kind of regardless, this kind of means that we might be able to significantly accelerate our memory access if the pattern of memory access is good. So if I want to read this entire block over here, if I access it in random order, then I'm going to have to basically query a number of times equal roughly to the length of my query. But if I go and I check the very first value, then I'm going to get all this entire burst section at once. And then if I go and check number 4, I'll get this burst section, the second burst section at once. And so I can basically get four times the throughput if I'm really clever about my memory accesses and only access just the bits I need from each burst section. So this is called memory coalescing. So if all the threads in a warp fall within the same burst, then basically the smart hardware and programming model will basically group those queries. Instead of querying 0, 1, 2, 3, it will group them and say, just give me 0, and then I will be able to read out all the 0, 1, 2, 3 at once from this kind of burst mode DRAM. So remember that a warp is 32 numbered threads. And so memory accesses from a warp happen together. And so when these warps are reading in to these kind of burst sections, there's optimizations that can be done so that you're getting all 4 bytes at once rather than getting one of them at a time individually. And so that will 4x the throughput that you have on your memory.
44:43 So these are very simple things, but they're actually very important. Imagine I'm going to do matrix multiplications, right? This is a core thing that you're going to have to do a ton. If you were to implement, let's say a neural network really from scratch in CUDA. In this case, imagine I'm going to read my matrices in one of two ways. I can read it by traversing the rows. So each thread is going to traverse the row. Or I can read it in column order. So each thread is going to go down a column. Turns out that this left one where you're going across different rows. So each thread is accessing a different-- oh, sorry, each thread is going through columns. This left model is going to be quite slow because the memory reads are not going to be coalesced. Whereas if you're going to this right side where each of the threads are going down, so they're incrementing in rows, then these memory reads will be coalesced. And so you can think about it for a moment why this is true. When I first looked at this diagram I was like, isn't it reversed? It's actually not. This is the correct one. And the way to think about this is let's say on this right-hand-side diagram over here, I'm going to have a thread that's trying to-- a series of threads that's trying to access left to right. So each thread is going to try to load the very first element. And then in the next time step I'm going to the load the element from this column, the second column and then the third column and the fourth column and so on. So if that happens, what happens at time step 1? At time step 1, my first thread loads at this point. And then the second thread loads at this point. And then this point and that point. So those can't be coalesced at all.
46:23 They're reading different burst sections. And so that means that I have to read this entire chunk of memory in order to perform any an operation. Instead, if I was going in the column direction, all the threads will be reading within this single burst section, and so only one memory read operation needs to be performed, and you get all of the memory at once. This is a very low level optimization, but this is very important. If your memory traversal order is all wrong, you will actually get much slower memory accesses than you really want. OK.
Trick 5: tiling, and where it goes wrong
46:55 So then that brings us to the very last and big one. And this is the idea of tiling. And tiling is this idea that you would like to group together memory accesses in order to minimize the amount of global memory access that we have to do. And so to explain this one, I'm going to try to go through this example of a matrix multiply. And hopefully, I'll be able to explain to you why a naive algorithm for doing matrix multiply is going to be very problematic. And then afterwards I'm going to give you a tiled version of the same idea, and hopefully you'll be able to see why that's going to reduce the number of global memory reads that you have to do. So let's start with this very simple matrix multiply algorithm. So I've got a matrix. I've got this M matrix on the left side. I've got my N matrix on the top. And in order to compute the matrix product, I'm going to have to traverse over the rows of M and the columns of N and then take the inner product and store that into this P matrix. Write the corresponding rows. And I've written out here each of the threads. The thread 0 corresponding to where they're storing their outputs and the access order in which they access each of the individual elements. Now notice here that what's going to happen is that the memory access here is not coalesced. Like, the row matrices here, these are going to be accessed in a non-coalesced order. And I have repeated memory accesses. So I've got M00 being accessed in the first thread, M00 accessed here, N10 being accessed
48:36 in two different threads. So these values are being read over and over from global memory into many different threads. So this is going to be potentially very slow. So there's a question of can we avoid having too many global memory reads and writes what I would ideally like to do. So let me explain the ideal outcome first, and then I'll explain the algorithm. The ideal outcome is that I would like to spend one chunk of time loading pieces from global memory to shared memory where things are fast. I want to do a ton of computation in shared memory, and then I want to be done with that piece of data. That's the ideal outcome. I've minimized my global memory accesses. So now how can I do in this matrix multiply world? So now what I'm going to do is I'm going to take my matrices, both the M matrix and the N matrix. And I'm going to cut them up right into tiles. So here I've cut this up into 2 by 2 tiles. So I've got a 2 by 2 m tile and a 2 by 2n tile. So I've got basically smaller submatrices within each of the matrix. And now imagine that my shared memory is big enough to be able to fit these submatrices within each of these SMs. So now this gives a very, very simple algorithm with which we can do computation. So what I'm going to do is I'm going to first load, let's say this M00 tile on the top left over here. And I'm going to also load my N00 tile into shared memory here. So now I have these partial sums that I can compute. I can take the row product of M00, M01 with N00, N10. And I can increment that into P0.
50:16 I can do the same with all of the different submatrices that I can fill out over here. Now then, once I'm completely done processing these two tiles, then I can load a new tile over here. And then I can repeat that computation with my M tile and my N2,0 tile loaded into shared memory. And then I can increment my partial sums in P. So now I've really consolidated and reduced the amount of global memory access I have to do. I load as much memory as I can at once into shared memory. I do all of my submatrix computations on that tile that I can. And then I move on to the next one. And of course, the other nice thing is that because I'm loading an entire tile, I can traverse these sub matrices in whatever order I want like column major or row major. And so I can coalesce all of the memory accesses whenever I'm loading a tile from global to shared memory. So there's wins all around here when we tile our accesses. So we can do a little bit of tiling math. So we've got let's say, a matrix A, a matrix B, and a matrix C. So let's say the full matrix, these are square matrices are of size N. And let's say I have a tile of size T. Oh, yes. Question. Previous slide of [INAUDIBLE] step three we're loading and M00 again? So in that case, I just wrote it for completeness. But M00 let's say is just stored in the shared memory. Let's just keep it cached. I won't load it again. Yeah, that's definitely just there for completeness. Not that you would actually discard and reload the matrix again.
51:53 That would be kind of insane. Cool. OK. And so we can do very simple tiling math to think about what's happening. So let's say I'm going to do an N by N matrix multiply right. So if I do a non-tiled matrix multiply, if I'm just going over rows and columns, then every input every time I process. It has to come from global memory. So each input is read N times from global memory. So each of these is read N times. If I do a tiled matrix multiply, well, the global reads are operating over tile. So I'm reading each input N over T times from global memory. And I'm reading T times within each tile. Of course, I'm doing matrix, matrix multiplies. So I can't reduce the total number of reads. I have to read all the matrix elements, but I can shift the reads into basically fast shared memory. So I do T times memory reads into shared memory and N over T times from global memory. And that's great because if we have a big shared memory that can store big tiles, that's a factor of T reduction in the total amount of data that has to come from global memory. So tiling can be really, really powerful of an idea when you're operating over matrices. And you can move things into shared memory. Tiling is quite complex. This is the source of many, many confusing things about GPU and matrix multiply performance. One thing that can happen. Once we start tiling things, you start asking things about discretization. So imagine I have a tile size of 128. That seems like a nice good round tile size. But then when I have a full matrix of 256 size, that's great.
53:35 That's a 2 by 2 tile, things load nicely. Now let's say I have a 257 size tile on the column side, now this is a bad time because I need to have six tiles in order to cover this matrix, and the two tiles on the right are very, very sparse. There's just not much stuff in there. And the problem with this is that each tile is going to be assigned to an SM. So each of these tiles is going to be a block. And each thread is going to be operating within each tile. So those two tiles on the right, they're not going to be doing very much at all. Those SMs are going to be basically be sitting idle. And if you were kind of compute capped, you would have wanted to more evenly distribute the load between SMs. So you have to basically optimize your tile sizes to try to avoid these kinds of scenarios. But in reality, there's a lot of complex things that go into setting the tile size. Remember, you have to coalesce your memory accesses. So you have to think carefully about that. You have to not exceed your shared memory size. So the tiles can't be too big. And you have to divide the matrix dimension, hopefully evenly or as close to evenly as possible. So you don't end up with this situation of an underutilized SM the very end here. Yes. So if you had, say, smaller sizes [INAUDIBLE] would GPUs do something like prefetching, where they can fetch the next tile beforehand? And so would that happen [INAUDIBLE]? Yeah. So you're asking about whether or not you can overlap memory reads and computation? And yeah, that's naturally done in GPUs.
55:17 They're always like trying to use the available bandwidth. As long as shared memory is available, they can go and put things into it. The issue is that whenever you're effectively utilizing your SMs, you're basically maxed out on your shared memory. That's like the bottleneck resource. And so there is no place to prefetch in some sense. Cool. OK. And the other thing that is very, very, we're getting into the weeds here, complex is the interaction between tiling and burst sections. So imagine I have a matrix layout that's this, where I have my nice burst sections, and each burst section lines up nicely with a tile. So to read this tile, all I have to do is to get four different burst sections. And I've gotten this entire tile. Now imagine what happens if I add one element extra and the way the matrix is laid out, my tile start-- my burst sections flow over. So now what's happening is when I load my tile, I'm going to load this first part. And that's really great. I get the entire first row as a burst section. Now in the second row, this actually belongs to two different burst sections. And so I have to do two reads in order to get this second row and so on and so forth. So I've essentially doubled the number of memory accesses because I've added a single extra element at the very end there that's kind of bumped up the alignment of my burst section and my align layout. And so basically, if tiles or your matrix sizes aren't multiples of your burst section, you can easily end up with situations like this
56:59 where the rows don't line up with the burst section, and you've doubled the amount of memory access that you have to do. And the way to get around this is you have to do padding to be able to get nice round matrix sizes so that your burst sections line up with the size of your tiles. So this is getting very into the weeds here. But if you really want to squeeze out all the performance from your matrix multiplies, these are the kinds of things you have to think about. And you will get bitten by this, if you're not thinking about it. And of course, I guess like things like torch.compile and all the CUDA optimizations for matrix multiplies, they're doing exactly the kinds of stuff that I just talked about. That's the way you get better performance. And so all of this matrix complexity ends up in situations like this where I'm reading out Andrej Karpathy's tweet here, but the most dramatic optimization to nanoGPT is to increase the vocab size from 50257 to 50304, which is the nearest multiple of 64, which gives you much, much higher occupancy.
The mystery solved: divisibility and wave quantization
58:06 Careful with your powers of 2. So that's a 25% speedup from adding how many? It's like 57, no, 47 dimensions to your vocab. That's like-- how does that happen? And so that kind of brings us back to the mystery. I was dragging you through all of the GPU details in the hopes that you'll have a full understanding of all the performance characteristics. But in some sense, the payoff is I now get to explain to you how this chart comes to be, and at the end you won't find matrix multiply performance to be so mysterious or scary at the end here. So the very first part is very, very simple. We understand compute intensity. This is exactly the roofline that I pointed out at the very beginning. So up until here, which is about 1536, there's just not enough matrix multiply work to do. Just loading the matrix and doing very basic IO that you have to do is becoming a bottleneck below this point. So throughput is going to fall through to the ground past this point. You just don't have enough memory bandwidth to support your compute units. Now on the right side here, in theory, if I draw the upper envelope, this is the kind of maximum achievable performance. So it's possible up here to saturate all of my compute units and get really great performance. But if you mess up your matrix sizing, you can end up in these kind of really weird places. And within each one of these you can end up in a weird trough. And so we're going to think a little bit about, why do you have all these different places you can end up? So the very first thing, this first line here, this
59:49 is a tiling alignment issue. So if you look at the multiples here, so I've now colored each of these lines based on the divisibility of the matrix size. And this is the size by which it's divisible. So if it's divisible by 32, then you're in good shape. You're in these purple dots up here. If you're divisible by 16, you're actually still up here. There's two colors. And then if you're green, your K equals 8. You're up here. If you're orange, your K equals 2. And if your K equals 1, you're all the way down here. If you're not divisible by any number, don't pick prime dimensions. You're not going to get very good throughput on your matrix multiplies. And a big part of this is going to be once you get to K equals 2 and K equals 1, you are basically forcing this situation where you can no longer read tiles in this nicely aligned way with your burst reads, and that's going to lead to some serious issues. So that's kind of a problem. But then, OK, so that's one part of the mystery. But I think another part of the mystery remains. Like, so within this orange line, I think if you zoom into here, you see this giant drop from this point all the way down to this point where you're just kind of wondering what happened here. How could I lose so much performance, increasing my dimension by 2? And so let's just look at these numbers. And it's just I think this is a fun puzzle. So I'm just going to walk you through the puzzle. This is going to happen when you transition from 1792 to 1790 I guess, 3 or 4 size, let's say 4 here, just so that it's a factor of 2 still.
61:27 Well, why does that happen? OK. Well, let's say that we're using a tile size of 256 by 128. That's a pretty natural size. As a fun fact, the matrix multiply units in these GPUs. There they are, naturally operating on matrices of roughly size 128. So 256 by 128 is a very nice tile size. So that means how many tiles are there? Well, there's 7 times 14 tiles because we're dividing the dimension of the matrix by the size of our tiles. That's a total of 98 different tiles. And if we increase this by one, well, we're going to have to round up each one of our coordinates. And so we're going to have a lot more tiles, 120 of them. So we've increased the number of tiles by quite a bit. Well, you know, what's going to happen is not only did we significantly increase the tiles and some of them have lower utilization, which is bad, but actually even worse. An A100 has 108 SMs. And if you go all the way back to the GPU execution model, SMs can execute in parallel and they're kind of the execution units. And so when you have 98 SMs, they all go and run. You can dispatch them all. All the SMs are running. You got great utilization. Once you go to 120 tiles, now you've got more tiles than SMs. So 108 of those will execute. And then you will go back and you'll say, all right, I've got some more SMs. At very, very low utilization, you're going to execute the remaining 12 and wait for those to complete. And that's going to be really bad. So if you look at your utilization, you've got good utilization for a while. You'll drop off a cliff, and then you'll finish up your job. So this is something called wave quantization. And so ideally your tile sizes are either much bigger than the number of SMs or they're not
63:08 this where you're just like barely over the SMs and you've caused this quantization error additionally. Cool. All right. I know this is low level details, but in many ways, I've been saying through many classes that language models and deep learning is attention to detail. And these kinds of attention to detail is the things that allow people to scale up LLMs to really, really large sizes and get great performance. So it's worth knowing, even if you're not a person that's going to do systems engineering. So what were the tricks right. Key ideas here. First one is you got to reduce the amount of memory accesses. So there's lots of ways to do it. You can do coalescing, so that you're not-- you can reuse reads that you're getting for free. You can do fusion so that you can fuse multiple operations together and avoid unnecessary reads and writes. You can move memory to shared memory. So even if you're going to do reads, they're going to be from much faster memory. And that's going to be tiling tricks that you can do. And then finally, you can trade memory for other resources that you do have. So you can trade it for compute, which is going to be recomputation. Or you can trade it for just numerical precision or stability, which is going to be quantization. So there's lots of bags of tricks that you have in order to get performance out. So there's lots of things you can do. You just have to be really mindful of the role that memory plays in the performance of a GPU. That's kind of the key thing to get the most out. Cool. Any questions on that before I move to the final part with FlashAttention?
FlashAttention, and the recap
64:44 OK, good. All right. So now I'm going to put it all together. I'm going to try to make it so that all the tricks that I taught you aren't these random disconnected facts about GPUs. They're kind part of the standard performance optimization toolkit. And FlashAttention and FlashAttention-2 will hopefully teach you how that all comes together to build one of the foundations, I guess, of modern high performance transformers. So FlashAttention, we know that it dramatically accelerates attention. And most of you probably know that that's done through some CUDA kernel magic. But maybe you don't know all the details. So what the paper says is, OK, so there's one part that's happening, which is do attention on a unoptimized PyTorch transformer implementation. If you fuse the kernel and you do some things, you can get significant speed ups. And from the paper, they say we apply two established techniques tiling and recomputation, to overcome the technical challenge of computing exact attention and subquadratic HBM accesses. So it's not subquadratic computation because you can't do that. You have to compute attention in general, but they're going to get subquadratic accesses to the high bandwidth or global memory. And so that's really the key. If your memory is the bottleneck, you want to make that not quadratic, so that at least you can pay for quadratic cost with your compute rather than with your memory. So just for a really quick recap at this point, you've implemented attention many, many times in many classes. So it's going to be three different matrix multiplies. You've got a K, Q and V with a softmax in between. So the matrix multiplies are pretty simple.
66:24 That can be done with tiling. I've showed you examples like that. And what's different about attention? Well, there's this softmax thing. That's going to be the real tricky bit. And then once we can deal with the softmax, all of the matrix multiply things I was talking about will just come into play. So the matrix multiply, as I said before, is exactly what I taught you. So if you look at the figure 1 from the FlashAttention paper, this is really just a simple tiled matrix multiply. You see the K matrix, the Q matrix, you see it cut up into small blocks. Small blocks of it are being copied to SRAM, they're being multiplied. And then they're being accumulate-- they're sent to the HBM where you do softmaxes and then you multiply with a V. So this is all just really simple in terms of the K, Q, V matrix multiply. But now we have to think about the softmax. What's going on with the softmax? So the key thing here is the softmax-- sorry, I'm going to roll back one step. So the issue with the softmax, what's the problem with the softmax? It's a global operation. The softmax in an attention operates row by row. You have to sum the entire row right to compute the sum normalizing term of the softmax. And that's very problematic. If I have tiles, ideally I want to do everything within the tiles. I don't ever want to have to write back to the big matrix. And so I need a softmax that can be computed online within each tile. I want to do as much computation within each tile as possible. So the key thing here is to use what's called the online softmax. And so what is that?
68:01 If you have a stream of values, normally the batch version of the softmax, you take all of your x1 through x of and you would exponentiate them, sum them, and you would divide them. That's what you would do in your normal softmax. And then you would maybe compute the maximum value, and you subtract that in order to be able to make this numerically stable. So this is the standard numerically stable softmax on the left side. So the online softmax, I've taken this from Milakov and Gimelshein in 2018. Well, you can realize that you can pull out via telescoping some kind of an argument, basically the current running normalizer term and the current top term of e to the xi minus max of xk. So what you're going to do is you're going to maintain your current max that you've seen over x1 through x of j, which is my current iteration. And then I'm also going to maintain this correction term. If my max updated, this is going to basically correct my max. And then I'm going to add my new term over here. So this d of j is going to track online the top term of this equation 2 over here. And then at the end, I can also then compute the normalizer and then get the normalized y of i that I want. This d of v is itself the normalization term that I need. So the key thing here is that this can be done online. I don't need the x1 through x of n up front. All I need is the stream of x1 through xn. And that's really key because I can now compute the softmax tile by tile. Within each tile, can run this algorithm, and that will let me compute the partial softmax for that tile. And then I can write back if I need to all the components
69:47 that I'm keeping track of. And that's all that I need in order to do this computation. So I never have to materialize the full n squared matrix in order to compute the softmax. And so that's basically it. Once you have that, you've put it all together and you can get the forward pass of FlashAttention. And if you go and look at the FlashAttention-2 paper, which is going to be a thing that we're going to ask you to implement. So you're going to be following through these steps here. You're going to see exactly this idea. So first, you're going to have your KQ matrix multiply. And this is going to be tiled. So these are little tile chunks. And they're going to be multiplied. And how am I going to compute the softmax. Well, I'm going to maintain a running value of these exponentiated sums. And then I'm going to keep incrementally updating it and correcting for the maximum terms. And by doing that, I can compute all of the necessary quantities tile by tile going from one tile to another, and then just multiply once again with tiles with v in the end. And that will give me my full softmax output. Yes. So we won't be able to compute that until we compute the QK multiplication across all of the tiles. So we do have to double back on each tile. So the question was you can't compute this until you are done with all the tiles. And so you have to double back on all the tiles. We'll talk about denominator sum until we've seen every tile. That's right. So you will have to-- before you can output your softmax, you will have to go through all the tiles. This is correct, but let's say I do all the tiles once.
71:31 I do all n square tiles. At that point, I have all the components that I need in order to directly output the softmax at that point. I don't have to do recomputation because I have the normalizer terms already. By going through each of these tiles, at the end of going through all these tiles, I've built up L3 or L of n, which is the sum of all of the exponentiated terms. So I already have that in my shared memory for this last tile. And then that allows me to exponentiate and divide and then return all of the components. So the backward pass I'm not going to cover. You can do recomputation tile by tile, which will allow you to avoid storing the softmax. Remember, I always want to avoid storing anything that's of size n squared. And so here I've been clever with the tiles so that I don't have to store any of the n squared components when I'm computing, for example, the softmax. But in the backwards pass, if I store the activations, that's already something that's n squared size. So I don't want to store my n squared activations. I'm going to have to recompute it on the fly tile by tile when I do the backwards pass. So that's a really key other trick that they do in order to make the backwards pass possible. But otherwise, it's fairly standard. It's the same thing as computing the gradients just tile by tile and doing that computation. So OK, that brings us to the end here. Hopefully you've seen how all of the pieces I talked about, about tiling and coalescing and recomputation, come together to give you FlashAttention and all these really cool things that make your transformers go much faster. So to recap for the whole lecture. Hardware is the thing that has really
73:16 powered all of the language models that we have today. And so if you really want to leverage your hardware, you have to understand the low level details. I think all the systems advances really engage with a lot of the concepts that I taught today. And the current GPU scaling, that plot is really the one you should remember, really, really incentivizes and encourages you to think about memory movement. The memory movement is the bottleneck in all of this. And so you don't want to just think about, oh, how do I reduce the number of FLOPs. That's important to really, you really have to think about, OK, how do I make my memory movements more efficient? And then finally, if you have to do a certain amount of computation, well, to optimize things, the way to do it is to optimize your data movement to be able to avoid as much movement from the high bandwidth memory or the global memory as possible. You want to reduce that and have everything in the very, very fast shared memory. And that leads to good performance on things like FlashAttention. Thanks, everyone.