CS336 // FIELD MAP
← deep dive
TRANSCRIPT · LECTURE 06Tatsunori Hashimoto · 80 min

Kernels, Triton

Cleaned auto-captions · timestamps open the video at that moment · caption errors corrected where the lecture script settles them, otherwise left as heard

Setup: what this lecture buys you in assignment 2

00:05 Today we're going to be going into details on make writing high performance code for GPUs. So part of assignment two is going to be you're going to have to you know do a bunch of profiling. You will have to write your own Triton kernel for FlashAttention-2. You will need to sort of make all of this stuff very high performance. And so in this lecture, we're going to kind of drill down a little bit and we're going to try to, you know, write some high performance code for standard components in a language model. so the plan for this lecture is we're going to just do a brief amount of review about GPU stuff. just to make sure you have once again the basic components of the GPUs that we need to understand in order to follow the rest of the lecture. and then I'm going to show you a bunch of sort of really basic things about benchmarking and profiling which will be helpful for both the assignment and in general if you want to write high performance PyTorch or deep learning code. And then we're going to basically write some kernels. we're going to write CUDA kernels in sort of C++. We will then do the same thing in Triton. And then lastly, we're going to, you know, do the easy but very good thing of using PyTorch's existing JIT compiler to have it optimized for us. And then we'll compare all of those and profile and benchmark things. And throughout we're going to really dig in deep. We're going to go down all the way to the PTX. So, so pretty close to the machine code to understand what you know the GPU is actually doing under the hood when we write all this code. and then hopefully we'll have time and I think we will we'll finish by writing sort of a fast Triton implementation of softmax at the very

01:45 End. Okay. So assignment one has come to a close. There's still a leaderboard. You can still submit and update things there. some of you may be using late days. So please finish up assignment one. and then assignment two is now out. And as I said before there's going to be you know a bunch of systems stuff that you're going to need to do. there's fun parts that you can do now involving GPU kernels and then next week we're going to talk about parallelism and that's going to be the other half of the assignment writing fast parallel code

GPU refresher: SMs, thread blocks, waves, arithmetic intensity

02:17 Like data parallelism and so on. So we will get to that next week. All right. So now remember how GPUs work, right? So when we have something like an A100 or an H100, we're going to have a whole bunch of SM streaming multiprocessors. within each SM is a large number of units that can do computation. we have INT32 units or FP32 units. and then each SM is going to launch a large number of threads, right? and we have the memory hierarchy. which is that we have DRAM or global memory which is big and slow. And then we've got caches that are much faster. and in fact, you know, there you see here there's this thing called a register file. This is very fast memory that each thread can access. And we're going to be making heavy use of these registers as we write high performance code for GPUs today. so the basic structure for the execution model is going to be we're going to have a collection of thread blocks and a block is going to be scheduled on a single SM. Right? So this is kind of the atomic unit that we're going to be thinking about especially when we write code in things like Triton. And then within each block there's going to be a whole bunch of threads and the threads are actually going to be the ones doing the computation. And so if you have a vector and you're going to be operating over elements of that vector, right, you're going to write code where each thread is going to go in and maybe operate over a few elements of that vector at once, right? And all the threads together will sort of process the vector completely. So why do we have these things called thread blocks, right? Why not just have threads and your big global context? Well, thread blocks can communicate with each other. There's shared memory

03:55 Kind of within the SM that's pretty fast, right? So when you need to do something like matrix multiplication, you're going to need to pass information from thread to thread. and within a thread block that's very fast across thread blocks or across these groups, it's going to be very expensive. So you any data that you need, you're going to want to keep within the same thread block or within the same sort of pile. and that's going to keep things very, very fast. and that's going to be as fast as sort of a L1 cache. And that's a great, you know, place to be. And so you can use this to synchronize across threads. but you can't you know for example synchronize across blocks you can't really control what's going to happen right and remember the thing that I mentioned last week there's this thing called waves right waves aren't sort of an inherent thing that you normally think about but for performance it is an important component so when we actually run these things the threads are grouped into consecutive blocks of 32 threads and that's a wave and that gets executed kind of all at once in an SM and so one thing that we would like to do is to make sure all the waves have an equal amount of computation. We can't always do that. but you know if we can we would like to do that right? So we want to make the number of thread blocks ideally divide the number of SMs and to make sure that each wave has an equal amount of work. So we're going to ideally have a lot more thread blocks than SMs. And we're going to try to make that happen as we write high performance code. Okay. And then the last concept and maybe amongst the most important concepts here is arithmetic intensity. we would like to keep arithmetic intensity high. we would like to have more flops than we have

05:30 Bytes of memory movement. and this is because you know if you remember the scaling plot from last lecture our compute scaling is much faster than memory scaling. So a lot of the time computations are going to end up being memory bound and we're not actually getting all of the work done right. So as a general rule you know matrix multiplication is compute-bound if we kind of do it cleverly. Everything else is going to be memory bound and we're going to try to cleverly reduce the amount of things that are memory bound or how badly things are memory bound. Okay. So that's our very brief sort of review of GPUs. Hopefully everyone remembers this. You still have a fresh sort of memory of the execution model. feel free to stop me and ask questions if any of you you know have sort of lingering doubts or questions about how this is all going to work. Yes. What was the function of warp? What was the function of sorry warp? A warp. a warp is essentially a group of threads that get executed together. And the reason why warps exist is that they reduce the amount of control machinery that's needed. because you're executing all these threads at the same time. you don't need a control thing for each thread. you need them for blocks of 32, right? And so you see, for example, there's a lot more compute units than there are sort of warp schedulers. and so you're able to do a lot more parallel work without worrying about control. And this is one of the trade-offs with CPUs, right? CPUs, a lot more sort of silicon area dedicated to control and branch prediction and things like this. Whereas for GPUs, much more emphasis on computation with simpler controls. Okay, so now we're going to get into sort of sort of newer content now.

Benchmark before you guess

07:09 And I think if there's one high-level thing to remember, it's if you want to write high performance code, you should remember to benchmark and profile your code. And that seems very obvious, but you know, I've seen a lot of things where, you know, students or people go in and they're like, well, I think this is the bottleneck, so I'm going to spend three hours optimizing it. And it turns out it wasn't the bottleneck at all. I'm sure it was fun, but that, you know, there were it was kind of time that was misallocated. And so if you actually use a high performance or very detailed profiler, you can kind of see exactly where your you know bottlenecks are and exactly what the machine is doing. And once you have that, you can go and spend your efforts in sort of the most important parts of your code execution. And so that's the high level thing I want to get across because some of the details about you know GPU execution and you know how you write a softmax kernel that's going to kind of change and maybe you even want to just rely on the torch compile you know auto-JIT thing. but the fact that you should profile isn't really going to change no matter what the tools are. So, I want you to sort of internalize that idea that you should be always profiling if you want to be writing high performance code. And really, you know, there's a limit to the theory. I think systems is part of this course that you can reason about pretty well. Architecture is somewhat hard to reason about and you can, you know, really think about sort of the roofline model and so on. But, you know, how fast does your matrix multiply? Well, maybe that depends on the library version or your hardware like which things are bottlenecking for what reason. There's all sorts of, you know, microcode things that you don't really fully know. And so, you have to in the end have to do end-to-end benchmarking whenever you're developing these things. Okay. So, I'm going to have an example computation. This is the

08:46 Simplest thing you know that we can run compared to all the things that you all are doing in your assignment one. but I'm going to run a very simple MLP. It's going to have 128 dimensions. It's going to have 16 layers. It's going to have some batch size and it's going to have five steps. I'm going to just do forwards and backwards for five different steps here. and just to make the code clear, it's it's something like this, right? I'm going to define a MLP model and we'll sort of I'll show you that in a moment here. and then I'll define, you know, a random Gaussian input and then I'll run it for five steps in that last case where I compute some forward and then I compute a backwards and then I return sort of the result which is just the mean of the output of my MLP, right? Not there's not even losses. It's so simple. It's just you run the MLP forward and I just average pool at the end, right? and then the MLP is just kind of the simplest thing you can also imagine here. It's just a bunch of linear layers stacked on top of each other. which is this bit and then you know I've got a GeLU in between, right? So this is just GeLU linear GeLU so on and so forth. Everything is nice and square, right? So hopefully this is a very simple MLP that you all feel pretty comfortable with. and then let's go back. Yes. Oh, sorry. I want to go back up to here. Okay, good. and so now I have this, you know, MLP code that I want to run. And now I'm going to do two things. I'm going to benchmark. So I'm going to do some timings. So I want to know how long does this function take to run? And then I'll do profiling, which is to go inside the function and ask, you know, where am I spending all of my time? So let's start with benchmarking, right? So benchmarking is just the measurement of wall clock time of performing these operations. and I'm only looking for the end-to-end

10:27 Execution time of, in this case, my MLP function. And you know, there are some subtleties to this, like you're sitting there and you're like, why am I being told how to invoke, I don't know, the time it function. but you do have to be a little bit careful about how you measure times. And I think, you know, if you're not paying attention, you will run into these pitfalls, when you do assignment, too. and so, what are we doing this for? We're going to compare implementations later. We're going to compare our Triton to our handwritten C++ to PyTorch's implementation and torch compile and we want to know was it worth it to write that CUDA kernel. and we'd also like to understand when I make my matrix multiplies bigger, how much slower does it get? Right? So we'd like to do some empirical benchmarking of those. So throughout this lecture I'm going to be using this benchmark function. and that's going to be sort of a wrapper function. I'll step through it. benchmark is going to do the following things, right? It's going to have a function that I want to benchmark, which is run. And then I'm going to do some number of warm-up iterations, and then I'll do some number of trials, right? and you might wonder, okay, so like what's this warm-up thing that we're doing here? Well, one thing that's really important is, you know, when you do when you first run your PyTorch code and let's say it dispatches something to the GPU, it might look very fast and transparent to you, but that very first time something is executed in the background, machine code is being compiled. you know, that code instruction might be being sent to the GPU. There's all sorts of things that happen to sort of initialize your code. and so you always want to do some warm-up iteration to make sure that you're not measuring sort of the startup speed. Instead, you want to measure kind of the steady state speed, right? If

12:04 You're running thousands and thousands of iterations, you know, what you're interested in is that part, not necessarily, you know, how fast can you, you know, do on the-fly compilation of your of your CUDA code, right? So, that's why we have warm-up, and you should always have a bit of warm-up. and then, another thing that's really important, and I'll get to this once we get to the profiler, is you want to call this thing called torch.cuda.synchronize(). Like, what is that? Well, the GPU and the CPU are basically two independent compute units in your in your computer, right? and they can basically run kind of independently. And so, their execution model is going to be this Python code that I have here. This lives on the CPU, right? And when I run something, it's going to dispatch a bunch of CUDA kernels, right, to the GPU. It says, "Please run these things for me, right?" And the GPU will go off and execute those things. And the CPU will actually go on and keep running, right? It doesn't wait for those CUDA executions to stop. And so that's great for writing high performance code, but you should hopefully see the immediate problem if you want to do benchmarking, right? If you're benchmarking and you've got this model where the GPU runs off in the side and your CPU is doing something different, you're actually not measuring the GPU execution time, right? so torch.cuda.synchronize() basically says, all right, let's make sure that the GPU and CPU are in the same state and there's sort of no queued things running and that we're we're kind of at the same point in terms of the code that's being executed. And now, so the GPU and CPU are kind of in the same state and I'm going to time it for real, right? and I'm going to time something for some number of times and I'm going to run the computation which in this case is the is the sleep command I'm going to do it three times and since I'm trying to sleep for 50 milliseconds that's the time that

13:43 I'm going to kind of get at the end right so I do time three times and of course here right I'm also calling torch.cuda.synchronize() at the end of run to make sure that the GPU and CPU states are the same. So, right, so the CPU is running ahead. It's going to wait for the GPU execution to actually finish here. and vice versa. and so now I sort of finished and then I'm going to average because you know each single measurement might be you know fluctuating because of things like thermal properties of the GPU and so you want to take multiple replicates take the mean and return that. That's our benchmarking code, right? Very simple,

What actually scales: matmul, steps, layers

14:17 But remember kind of the two important pieces here, right? Always do a warm-up. Make sure to call cudaDeviceSynchronize. if you do those, it's very simple. If you get forget to do those, you'll get pretty crazy numbers like you'll get that your big matrix multiply finished instantly, which is definitely not true, right? Okay. So, now we can do some benchmarking of matrix multiplies. I'm going to walk through some of these. they're just putting numbers to things that we already know, but I want to, you know, just walk through it and make sure we're on the same page here, right? So, I ran this on the on the class H100s. I have GPUs. I'm going to do matrix multiplies over these sizes. and then I'm going to go and collect a whole bunch of matrix multiply timings for each of these dimensions stepping through kind of this benchmark result. And so, we kind of see, you know, as we expect, right, super linear scaling of our runtimes as we increase the matrix size. Of course, at the smallest sizes like 1024 and 2048, we actually see that the times don't grow at all because there's constant factor overhead in just doing these matrix multiplies like these numbers have to get shipped from the CPU to the GPU. you know, there's overhead in like launching the kernel. and so it's not the case that you know it's super linear all the way to zero. but once the matrices get big enough, we see exactly the kind of scaling that we expect to see with our matrix multiplies, right? Okay. So, hopefully straightforward. Now, let's try to benchmark, our MLP. So, what are we going to do? We're going to make our MLP bigger. We're going to have 256 dimensions. We're going to have four layers, batch size of 256, take two steps. and so, what's the time that

15:56 It takes to do that? Well, it's going to take 6.2 seconds to do that. And now I could do some basic things. I can scale the number of steps from two to five and I can benchmark all of those and I'll get 2 3 four and then five steps. And unlike in the in the matrix multiply case, right, if I'm scaling the number of steps, so the number of forward and backward passes on my MLP, right? What do I expect the runtime to behave like? Well, I expect sort of linear scaling, right? And that's kind of what we see. there's about five seconds per MLP execution and we see it's about n times five for the runtime of kind of the end-to-end object here right okay let me see if I can reset the thing that's being monitored here oh nope I can't okay I'm going to zoom out a little bit sorry about that okay now we can also scale the number of layers from 2 three four to five and what does that give us well it gives us you know increasing run times once again linear in the number of layers, right? This time once again one layer takes about 5 seconds a little bit less than that and so we get about four times actually four times the number of layers and linear scaling sort of shows up again. Unsurprising, right? So both steps and layers obviously have linear relationships with the runtime and that is exactly kind of what we end up seeing at the end here. I'm going to skip the batch size thing because this is getting a little bit unwieldy in terms of the amount of things that are being tracked here. Okay. All right. So, that's the end of this benchmarking bit. We can kind of make this nice function that does a little bit of warm-up, does cudaDeviceSynchronize, and we can measure the

Profiling: seeing under the PyTorch surface

17:37 Runtime of anything that we want. And this is good, and you should do this all the time in your code, right? You can measure how long it takes for your new fancy architecture to run. But then I think if you want to fix some problems, benchmarking is a very coarse grain tool. It tells you that your code is slow, but it doesn't tell you where the time is being spent. And so what we would like to do, is instead do, profiling. and so this is going to be a much more fine grained object that we're going to want to do. and so profiling is really nice because it not only helps you see what where the time is being spent, which functions, but you know, when you look at what you're calling, usually you interact with the PyTorch interface, right? Like the parts of PyTorch that you call, but beneath PyTorch, there's this whole universe of CUDA stuff that's being called. And when you run a profiler, you can actually see all the way to the low-level calls what is actually being called. And so you can get a much nicer intuition for how the program is actually being executed on the hardware. And so we'll step through profiling a few simple functions and then get a little bit of intuition about what is happening. And so one of the things that is nice is that if you want basic profiling PyTorch has a very nice kind of built-in profiler that you can use. and this will allow you to not leave the Python PyTorch world and get some fairly reasonable looking outputs. And so I've profiled some functions here and you can kind of see the output of this as well. and so you know I've taken the sleep example from before. and here is you know the sleep function and when we profile the sleep function the profile function

19:16 Looks something like this. you know, I have a warm-up again. I have torch.cuda.synchronize(). and then I call the profiler and I'm tracking both CPU and the GPU times. and then, you know, I run something and then I synchronize again and I print out the average table across all the time. Okay. So, I go back now. So, now I'm going to profile the sleep function. and if we look at, you know, what's happening what happens here? Well, 100% of the time is being spent on something called cudaDeviceSynchronize. because there's no GPU work being done. This is just kind of a no-op. you know, it's kind of a silly thing to be profiling. And so now let's look at something kind of non-trivial, right? So let's look at this basic operation here of adding two matrices, right? So I defined a add function that takes in an A and a B and adds them together. and this is a helper function that instantiates two random Gaussian matrices and then invokes you know whatever is the in the operation argument. So this is adding two 2048 size matrices together. Okay. So now I'm going to profile this and I'm going to call the profiler and I'll get back something that looks like this block over here. Right? So this is what I get back. and I'm going to have to zoom back out because this is not going to be all righty. Okay. is this visible from the back? Can someone give me a thumbs up if it's visible from the back? And Okay, good, good, good. Or thumbs down if it's not. All right, so when we when we call the add function in Python, right, this is kind of all that we interact with this add function a plus b, right? That's all we think about. But actually underneath here, the underneath the iceberg so to

20:54 Speak, there's a lot more that happens. So this gets dispatched to the GPU and first there's this thing called ATen, which is the C sort of interface for PyTorch. And so this wrapper gets called and it says okay I'm going to add some numbers right this is what's being called that's the outer wrapper and then that dispatches to a particular kernel called vectorized_elementwise_kernel<4, at::native::CUDAFunctor_add dot right and this is the thing that's actually doing the adding and then there's this also other thing called cudaLaunchKernel that's taking some time and this is actually you know the CPU is taking the command and sending it over to the GPU that's the kernel launch and that takes some time and then finally you know the cudaDeviceSynchronize fires we're waiting for the GPU to finish and send things back to us and that also takes some time right the mere act of having a synchronization barrier is going to cost us some time and so we basically have you know the time total in the end here 1.4 milliseconds on the CPU and 17 microsconds on the CUDA. Right? So, so they're really fast on the GPU, slower on the CPU. And if we're looking at the CPU time that's being spent, which is the self CPU time, we see that kind of the C++ interface or the C interface is actually the thing that's costing us a whole bunch of CPU time. And there's sort of overhead to doing anything where we're sending stuff over to the GPU. So, that's the ad function. and we see you know what's happening under the hood. Same story here if I want to do a matrix multiply. So I'm doing you know a multiplied by b. So this is a matrix multiply of a and b you know I'm doing 2048 matrices once again. And then I do profiling. now this time I see you know ATen matmul. So this is saying

22:36 Like this is the lower level interface to do matrix multiplies. and this is going to dispatch the CUTLASS, which is NVIDIA's sort of high performance matrix multiply CUDA library. And then it's dispatching to a very particular CUTLASS kernel, which is going to have some tile size. the names are truncated here. I'll show you a more detailed version in a minute. you know, there this is basically pointing towards a very particular set of like tile sizes, and the number of blocks and so on. And so this thing is parameterized. and that's actually doing the matrix multiply. And once again we see the same two things at the bottom here, you know, the kernel launch and the synchronization of CUDA devices. and you can sort of see once again the CPU time CUDA time split. And we're spending way more time in CUDA because you know matrix multiplies do take more time than just adding two vectors. Okay. any questions so far? I can I can pause for a moment here. I think I've just been going sort of very quickly and on my own through the profiler. So if anyone has questions I can I can stop for a moment. If not I can keep going. Okay. Oh yes. In this case our time is greater than our CPU time but we did have a barrier that like said to for the CPU to wait for it to synchronize and so by that shouldn't the CPU time always be at least the same time? Counting the time. Yeah. I don't I don't think this counts the time. Cool. Oh yes. Sorry. there's too much there. is there any particular reason why like when we switch from adding to matt the CPU time went down? is there a reason why when we go from adding to matmul the CPU time goes down? That I

24:17 Am not sure to be entirely honest. Yes. Is there time compared to like running it? Is there overhead in the profiler that can distort things compared to running it in the real world? yes there is overhead in the profiler. like the barriers will do that. I'll show you a more advanced profiler from NVIDIA and you can add things like annotations that will also slightly distort the timings but not by much. the really large scale things that you see aren't going to be really distorted by the profiler. so if you're looking at like micro timings, yes, probably. But a lot of the things that we care about in the class, no. Yes. Just to make sure I'm interpreting this correctly. So is that like for the ad case is the 98% CPU being utilized over the time period that it's like the millisecond time period. That's right. Yeah. So this is the percentage of time as you can see that the actual millisecond time that ATen ad was actually executing in some capacity on the CPU. I don't think the CPU% of what the CPU is doing. Yeah, that's right. This is the time that the CPU is active, not percentage utilization if that's Yeah. So, this is not like the total amount of CPU flops or something. This is a total percentage of time that the CPU is doing something. Yes. Okay. Cool. All right. here's another example of a matmul. so this is a different dimensionality, right? So, this is a I'm multiplying 128 dimensional matrix

25:55 Here. so 128 by 128, much smaller. and you'll actually see that now it's actually directly executing sort of this different command. It's executing xmma GEMM. GEMM is the a matrix multiply type and this is float 32 float 32. You can kind of see from the naming of this kernel what's actually happening here which is that this is a tiled matrix multiply of some kind and it's not sort of going through CUTLASS. It's executing this particular command directly. And so for a small matrix multiply, you know, you see that it's dispatching to a different kernel. Now, so you can kind of see kind of the complexity of matrix multiply when we're operating at this high level abstraction, we just think of matrix multiply as a single thing, right? We call like a at b and we're done. But underneath the hood, depending on the dimensionality that you have, depending on the hardware that you have, it will actually dispatch to very different matrix multiply sort of primitives under the hood. And that will actually manifest in very different sort of performance characteristics. And so one fun tip is torch compile which I will talk about later actually has an option to sort of microbenchmark the matrix multiply performance on your hardware and then it will actually then pick the highest performing matrix multiply subroutines for your for your model which you know in the past I found you know gives you like 10% speed ups for free. It's very cool that like optimizing for these things actually gives you free gains out in the real world. Okay. so that's another matmul example. and so the cool thing about the profiler compared to the just the

Composite ops: cdist decomposes, GeLU and softmax do not

27:33 Raw benchmarking is we can now kind of see which CUDA kernels are being called. we can see that you know different sizes of matrices lead to different CUDA kernels. and we see you know cutlass_80_simt right is a is a diff is this CUTLASS linear algebra library and it tells us things like the t tile size. So, so far these operations are very boring in a way like matrix multiplies and adds they're basically one to one. You have a you know operation on the CPU side, it translates to a GPU operation and it just gets shipped over, right? So there's just a single operation in all of these that does anything on the GPU. So I want to look at some more complicated operations two more of these that have sort of more compound behavior. So what I want to do now is I want to do I want to look at this operation called torch.cdist and this is computing you know for two sets of matrices the pair-wise Euclidean distance between two sets of vectors right so this is going to be a big distance matrix computation between a's and b's that I want so that's c dist and so this is obviously a much more complicated operation if you want to compute Euclidean distances you're going to need to compute dotproducts you're going to need to compute square roots and we're going to see that once we compute cdist so now here is the is the profiled output of cdist. so we see that this torch you know python command does map in the in the c interface to some sort of lower level cdist. So this is ATen cdist which then maps to ATen Euclidean distance. and then

29:12 This will decompose into a whole bunch of things like ATen mm mole ATen pow and then sum because these are all primitives that you're going to need in order to actually to compute the Euclidean distances between all of your vectors and when you for each one of these like matrix multiplies and concatenation and taking the powers you have a corresponding cuda command that is being called here you know we have gmm which become we've become familiar with So this is a matrix multiply. It's taking 78% of our compute or our compute time on the GPU. we've got you know copies and sort of concatenation of arrays. This takes 6% of the execution time and then this sort of vectorized_elementwise_kernel which is taking the power takes 5% of the GPU time and 3% goes to the sum. So now we get this very nice low-level breakdown of where, you know, my GPU is spending all of its time. and from this, you know, I can get some sense of where maybe I should spend my time optimizing. you know, maybe I think I can optimize my matrix multiply. That would be great because that's 70 plus% of the time spent in the GPU. The final example the final two examples, sorry, that I want to talk about is GeLU and softmax. So these will be our running Oh, sorry, there's a question. What's the too wild. okay. So, I will maybe answer that question in a in a few minutes because there's a cooler profiler that shows you a much nicer picture and so I can gesticulate here, but I think it'll be better to show that with pictures.

30:51 Okay. So, I'm going to talk about now the GeLU and the softmax. so the GeLU is going to be our running example throughout the class. So, this is a nonlinearity. If you remember, it's the Gaussian error linear unit. and that's going to be a product of a tanh and a exponential if I remember right. and so we're going to have you know all sorts of operations. So we're going to add a and b and then we're going to call GeLU sort of simulating the linear plus nonlinear structure that we might have in our MLP. And so we see once again basically the same sort of mapping. we see ATen add corresponding to a plus b and then we have the cuda equivalent and then we have actually a GeLU function implemented in cuda which is all the way down here and that takes about 33% of the compute okay fairly reasonable and then we have once again the softmax I won't go through all of these in sort of gory detail since you know they all start to look the same after a while but the thing to really point out that I think is cool is that a lot of these really core primitives like softmax and GeLU there's kernels written for them, right? So, it's not like the GPU is executing the basic primitives. There's sort of a fused operator that computes all of this. So, there's no back and forth between CPU and GPU for all of these. So, okay. I mentioned before that I was going to sort of answer this question of what the CPU was doing. and so let's think about something a little more sophisticated, right? I had the MLP example that I started with for benchmarking. and I would, let's say, like to optimize that MLP, make it

Nsight Systems: the CPU runs a whole step ahead

32:28 Run really fast. So how can we do that? Well, ideally we would sort of profile this in a nice sort of fine grained way. So if we use the torch profiler, this is kind of what we would get. if you remember the MLP, there's you know stack linear layers. There's a forward and a backward. and you see roughly, you know, there's this backward thing that's happening. There's a matrix multiply. There's linear. and then there's accumulate grad operation for the backward. and here's the matrix multiply kernel. And then there's only 10 things that can fit here. So I think this gets cut off at a certain point. But this is nice. It does tell you that most of the time is being spent in the matmuls. but you do kind of wonder like where does all the rest of the time go and why does only 31% of my time stay here and where's the 60% here? It's a ATen mm but there's no corresponding kernel. Right? This is a little bit mysterious and for something that's very complex module this is not a very good visualization and so for that I think we have to actually get out a real sort of grown-up profiler and you will have to you know or we will ask you to look at this thing which is NVIDIA's Nsight Systems and this is the kind of NVIDIA's sort of detailed way of looking at GPU behavior and performance And so we will actually kind of see exactly what is happening as we run this MLP. So actually in the back can you see I don't know this tiny text over here. Thumbs up. Okay. All right. If you can see it then I'm not going to zoom in but it does it does seem small

34:05 Even from here. all right. So basically if we look here we see several different things. We see CUDA HW over here and then we see threads. and so this top half, this CUDA part, this is what the GPU is kind of doing. And then in this threads part, we see kind of what the CPU is doing. And I can also pull up the code, I think. Yes. the code here, when I profiled it, I've added a few annotations. okay, this one I zoom in for sure. okay. Let's, excellent. All right. so I've annotated the code with this set of things that says let's see NVTX which basically annotates my code with annotate with markers. So when the profiler comes in here it will know that this piece of code belongs to a block called define model. And for example the this part that says step range push and range pop. this range here from line 77 to line 55 should be annotated with something that says step underscore step. Okay, so I've added all these annotations in my code before calling my profiler. And so let's go back here. So now if we go to this line that says NVTX, we can kind of see define model which is the thing that I wrapped my model construction call. And then I see step zero, step one, step two, step three, step four, step five. So each step is now nicely annotated in this profiler and we can kind of see all of the things that the model is doing as we as it goes along and I'll start on this side. One thing we see is that this

35:47 Piece of code it doesn't do very much work. It takes only 14 seconds. So actually most of the time for the profiler is spent on overhead. So the part up until roughly here is you know things like just loading the libraries and that takes a long time. It takes apparently 7.5 seconds. just initialize everything and then on at least on the GPU at 7.5 seconds or so into the program it starts actually building the model and you see here on the memory footprint you know this is the place where now memory is being sort of allocated and on the GPU memory the memory usage starts to grow right now the model is now constructed at this point and then step zero is where sort of the action starts to happen and so you were asking earlier what's happening between the CPU and sort of GPU. And so how the execution model of this works is here is sort of step zero on the CPU. And I'm starting right here and here's the forward pass and this is layer zero. So let's just kind of think through what's happening. as I said before when you first encounter or when you first call a piece of code in PyTorch it doesn't just directly execute. it will actually do things like you know on the fly compile things and so you know this thing like runtime triggered module loading is sort of overhead work that's being done in order to just initialize the layer and the computation and move sort of various bits of code into the GPU. So this takes a long time. and then after this layer zero is done now if I look at sort of any slice here let's sort of zoom in to selection we'll see that each of these

37:24 Layers is really quick and what happens here is when I highlight this layer one over here on the CPU side notice that's not where layer 1 is on the GPU side right so as I said before the CPU and GPU are kind of two different execution devices so I start at layer zero I'm done with layer zero I start layer one. Now, the CPU is actually just sending all of the sort of CUDA commands the CUDA kernels it's launching all the CUDA kernels already to the GPU at this point, right? So, when the CPU is saying, I'm doing layer one, what it's actually doing is it's queuing commands into the GPU. It says, "Now run this thing next. Run this thing next. Run this thing next." Right? and so the CPU is running way ahead of the GPU. And by the time layer 1 starts executing on the GPU, actually, we're already at layer 9 on the CPU, right? Right? The CPU is running way ahead and there's basically a queue that the CPU maintains where it's sending a fixed number of kernel CUDA kernels to the GPU. And so once you hit that Q depth, it's going to sort of stop running ahead. But until that point, it's just going to keep going and going and going as far as it can, right? and in this case, this does become I'm gonna zoom out again. okay, undo the zoom. There we go. in this case, this kind of gets a little extreme because if I zoom out once more, notice how, you know, in these steps, I'm running way ahead. Like the step zero is here, step two is here. This was step one, which basically took no time at all. step two is here.

Why a print statement reshapes your GPU timeline

38:59 So, it's the CPU is basically running one entire step forward and backward ahead of the GPU. one interesting thing that you might do is if you're writing, you know, various code for training a language model. One normal thing that you might do is let's go back to the code. I might do something like print, you know, my losses in between iterations. this seems like it should have no effect on what the GPU is doing, right? You're like, well, it's a print statement. How much could it could it do? if you think about it for a moment, this will have big impacts on the execution layout on the GPU because in order to print this statement, right, this print statement happens on the CPU and the CPU needs to get the loss. That means it needs to wait for the GPU to compute that loss. And so let's look at what happens. So here, you know, as I said, you know, step four on the CPU happens way before the GPU equivalent. Now, let's switch back. Now, this is the version that I profiled where it has the print statement, right? And then now I sort of zoom into selection here. Now see how step one and step two are basically kind of synchronized now, right? Because I have to wait for the loss to get computed. And you look at this and you say, "Oh, but it's still a little offset, right? Like step two, step one isn't exactly aligned with each other." So now let's kind of zoom back in and see, okay, what happened to step one on the CPU? Well, basically the end point of step one on the CPU is also kind of where the optimizer step starts, right? So by the time that forward is done, sorry, this CUDA stream synchronizes the thing. So this cudaStreamSynchronize command on the CPU, this is basically saying I'm just

40:37 Waiting for the GPU because I can't run ahead. I'm waiting for this loss to be computed and to be spent sent back to me, right? So this is kind of a dummy operation where it's saying CPU waits, waits, waits, waits, waits, waits, waits. well, the backward step is done. So now I can print the loss. I've printed the loss. Okay, now the CPU can start running ahead. And it does run ahead and starts sending step two stuff now. And then well, once this hits here, it's sort of run out of commands. It's waiting for the loss again. cudaDeviceSynchronize. Wait, wait, wait, wait, wait. Backward step is done. Now I can print the loss. Now I run ahead again. Right? So in this case, you know, the GPU is still essentially full utilization in both cases. But in extreme cases where let's say you're printing tons of stuff all the time, actually you're going to introduce a CPU bottleneck, right? Because the GPU has to the CPU has to keep waiting for the GPU and it can't launch the kernels sort of ahead of time. So that's kind of a really cool thing that you can see with the profiler sort of this CPU versus GPU and they're actually different devices that communicate to each other. It's not at this single unified object and you wouldn't see that unless you started to look at some of these like more advanced profilers. any question about that sort of set of things? Cool. Okay. and the other thing that I want to kind of show you is you know the profiler thing that I was playing with before. You can also generate very similar views in Nsight Systems as well where you sort of select some range of things that you want to let's let's do a warm-up. I said we should so we should exclude the first couple of steps. So we'll start at step three and we'll we'll measure some

42:13 Steps. sort of in this range we could take the kernels. This is what's doing the computation. And you can see that there's actually many different kinds of matrix multiply. This is one matrix multiply kernel. This is a different matrix multiply kernel. There's a different sort of like vectorized_elementwise_kernel. and all of these are taking different amounts of computation. And we can take this and we can say oh show me in the events view all the things that are happening. and I can also see sort of the stats view all of the time that it takes. Wait, let's see. We want we want the average time. No, we want sorry the CUDA kernel execution summary. Yeah, we want the total duration of the kernels and so we can see which kernels are taking the most time and aggregate across these views. So this is actually a very powerful tool that can give you both like the aggregate view of what's slow and what's fast as well as individual kernels that are being launched and when they're launched and where the CPU commands for that came from. and I guess one final side note here is this is one of the reasons why you know it doesn't matter that we're programming in Python and Python's not a very high performance language, right? Because the CPU is never the bottleneck because the CPU can run ahead and sort of cue commands into the GPU. and so this sort of detaching or like this disconnecting aspect between the GPU and the CPU is one of the key reasons why we can use this nice high-level programming language and yet still get sort of full utilization out of sort of our

43:53 GPUs. Cool. Okay. Any questions before I sort of switch back to this because I'm going to leave Nsight Systems sort of forever for this lecture at this point. Cool. Yeah, but you'll get to play with it in assignment two, and I think you'll appreciate it because it gives you like a really interesting view into what your hardware is actually doing to make these like language models train. So, okay, that was benchmarking and profiling. Now, you have all the tools you need to be able to do sort of performance things. and now we're going to write some

Kernel fusion, and the GeLU gap it opens

44:25 Kernels in the remaining time. So, remember kernel fusion, right? So, this was the image that I showed you in lecture, right? there's a little factory. Every time I need to do an operation, I need to ship it from the warehouse to the factory and back. And so if I, you know, naively do a bunch of operations in sequence without thinking about it, I'm paying for a lot of sort of shipping cost back and forth from the warehouse, what I should do is have one factory that does all the operations at once. So I do not pay for this cost multiple times, right? That's very important. So now we're going to do GeLU. And we're going to write a kernel for GeLU. And I'm going to write that kernel in several different ways. And we're going to look at the performance impact of doing that. and so we have the PyTorch implementation of GeLU. And that looks just like this. torchn functional GeLU. and I invoke approximate equals tanh because I want this to exactly match the naive thing that I'm going to do next. So this is not going to be, you know, actually multiplying by the CDF of the Gaussian. it's going to be some approximation to that's easier to compute. Okay, so that's the PyTorch GeLU. And now I'm going to do the dumb thing, right? I'm you're going to look at this code and say this is going to be low performance. I'm going to go in and in PyTorch I'm going to write GeLU as 0.5 * X * 1 + tanh(sqrt(2/pi) * (x + 0.044715 * x cubed))). Right? magic formula, but this is a good approximation to the GeLU. can you can look it up or convince yourself this is true. but if you do this you see that there's a lot of operations

46:03 That happen right there's like a tanh there's a x cubed there's multiplication by a constant in addition and multiplication by 0.5 and x if this involves you know multiple different CUDA kernels this is probably going to be slow right that should be our intuition at this point from fusion so let's see if that's true okay so these two are the same you can see at the top left they compute the exact same numbers and you know we can systematically check this on random Gaussian. And now let's sort of benchmark the two. Okay, so the manual time is 8.1 seconds for a really big GeLU. and PyTorch time is 1.1, right? milliseconds, sorry. and the fuse version is going to be significantly faster. In fact, eight times faster. Wow. you know, big difference from writing a simple kernel. of course your matmuls are probably still going to be the bottleneck, but it would be really cool if we could go from that 8 milliseconds to that 1 millisecond, right? That would feel very satisfying. So, we're going to try to get close to that 1.1 millisecond in the next few parts of the lecture. So, now let's look at the what's happening under the hood. I don't need to look at Nsight Systems because all I really want to know is some very high level stuff for the manual GeLU. you know, kind of just like I said, it's going to do a whole bunch of operations. It's going to do a bunch of multiplications. It's vectorized, but it's a bunch of, you know, CUDA kernels being launched here. and notice on the right, this CUDA kernel gets called three times because we have a whole bunch of multiplications floating around here. we've also got, you know, addition. We've got a tanh. and each one of these is probably kind of slow and in the end, you know, we're incurring fairly large overhead doing this. now let's do the same thing, sorry, with the PyTorch GeLU. And this

47:47 Is this is really great. There's a single CUDA kernel launch. It happens once and it just processes the whole thing. This is what we'd like to see. and of course this is very fast because it's just a single CUDA kernel, right? So this is really nice and we would like to you know somehow get to the CUDA kernel. And so the first thing you might think of depending on how much you know about writing GPU efficient code is all right the PyTorch people must have written this in the lowest level language possible. So we're going to do the same

Writing GeLU as a CUDA kernel

48:18 Thing. We're going to go to not the lowest level possible but we're going to go to the C++ API and we're going to write the CUDA kernel in C++ right? So let's open it up and write our own CUDA kernel. So how is that going to work? Okay, so we have gone in and sort of created a C++ version of the whole thing. So CUDA, you know, when we say CUDA is actually the C++ API for interfacing with and programming GPUs. And just like sort of the logical model of a GPU that we describe, you know, we're going to write some sort of function f. and then when we sort of invoke this CUDA kernel, it's going to automatically call F on all the elements of a vector or a matrix. and then we will get to parallel compute everything that we want. as nomenclature we're going to have a grid which is a collection of thread blocks. So think of this as I have a task. I'm going to cut it up into pieces. and there's going to be a number of blocks. This is the you know in a 2D grid for example. there's going to be sort of a row coordinate and then there's going to be a column coordinate. And this will be very useful if you're working with matrices. And then there will be the size of each of these blocks like you know how big are these in terms of the number of thread blocks. So this is the dimension of the blocks. and then there's a collection of threads within these blocks and this is the coordinate that for example one thread block lives in and then each thread is within each block. Right? So there's sort of hierarchical structure here. There's a grid and then there's a thread inside a grid. Right? And then we're going to basically each function is going to take in three things. It's

49:53 Going to take the blockIdx like which thread block do I belong to which what's kind of the block dimensions and then what is the index that I am like my thread index and with these I can kind of know which coordinate that I am in the matrix or the vector and then I can sort of decide what logic that I want. one sort of last thing before we go through the actual C++ code is, you know, whenever you're you're trying to debug CUDA, you want to launch with CUDA_LAUNCH_BLOCKING=1. This will allow you to actually debug your CUDA kernel. It will give you sort of error messages back at a at a cost in terms of the runtime. if you don't do that, you are going to have a bad time if you're writing CUDA code and needing to debug. So, okay. here is my GeLU code and let's go through it kind of piece by piece and then I'll talk about what all the pieces are doing. this will probably take the longest out of the things that we're going to walk through. other than the machine code. and once you understand this, you should be able to understand all the other pieces. So, we'll go through this a little slowly. so there's two parts of this code. So, the first part, this gelu_kernel piece up here, this is the actual kernel. This does the computation, right? This is going to get sent to the GPU. It's going to do the computation and then it will return the results. This piece, the GeLU function here, this is a wrapper, right? This is lives on the CPU. It's going to orchestrate the launch of the kernel which is actually going to go out and live in the GPU, right? so maybe we can start with kind of this sort of wrapper piece, this GeLU function first,

51:30 Right? So we're always going to check two things. basically in the Triton or the CUDA code, we're always going to check. Oh, sorry. There's a question back there. Okay. Sorry, that's my bad. Okay, let me zoom in. That is an easy fix. but I needed to know that you can't see. Okay, good. all right. Is this good? Okay, excellent. okay. So, we're going to start with the gelu function. And there's two things that we're we're always going to need to do. The first one is to make sure that X lives in like the GPU device, like the CUDA tensor of some kind, right? If it's not well that's going to be a problem we're not going to be able to do anything on the GPU. The second thing which is maybe less obvious is that we want to check to make sure X is contiguous. What that means is it lives in a contiguous block of memory because when we index into X, we're going to do a whole bunch of indexing arithmetic and we're going to assume that X lives in a block of memory, right? And if it doesn't, it's just going to be, you know, basically impossible to do this with any level of generality. and so when we compute the GeLU, right, we take in an input X and we're going to output a Y, right? And so we need to allocate a output. So torch::Tensor y = torch::empty_like(x). This is just saying well give me sort of a output tensor space or a pointer to a output tensor that is just like the dimension of x and notice that I'm not calling zeros. This will save on extra operations. I don't need to zero out these y's because I'm going to write into them anyway, right? So this is a minor but you might as well do it optimization. And then basically in

53:11 All the code that we write, we're going to need to figure out the grid, right? So what's the total number of elements that I have? What's the size of each block? The number of threads that I have in each block. And then how many blocks total do I have? And when I need to figure out the number of blocks, I'm going to, you know, call cdiv, which is going to be essentially take the ratio of num_elements to block size and then take the ceiling, right? Because I need to round up to make sure that very last set of elements that sort of isn't divisible by block size still gets computed, right? So I take the ceiling rather than the floor. And then this is all very simple bookkeeping stuff. And then I say all right launch the kernel. you know the gelu_kernel gets launched. and this sort of angle brackets is saying this is kind of the with the given number of blocks and the and the size of each block. And this is going to be passed into sort of the kernel command. And then I'm going to pass in the pointers to x's and y's, right? I'm not actually going to pass the values of x's and y's and the total number of elements. And I need this to compute sort of essentially the boundary conditions of my kernel. So now let's go to the actual kernel itself. Right? So I have __global__ void gelu_kernel and I get in pointers for in and out and I have number of elements items. and this keyword global the website sorry the rendering here has mangled it a bit a little bit but you should think of this as underscore global and this is a keyword that distinguishes it as a as a CUDA kernel function. And so what am I doing? Well, you know, this thread is actually supposed to operate on a single

54:47 Element I, right? but I don't get I as input. Like the code doesn't actually tell me you're in a vector in coordinate I. So I need to compute where I am. And how I'm how am I going to do that? It's going to be I take my blockIdx, right? I only have one dimension. So it's blockIdx.x. So just the first coordinate. and then multiply it by the size of each block. The blockDim.x and this tells me, you know, basically the starting point within my current block. And then now I add in threadIdx. So, you know, I know where the start of my current block is and I add in the offset to where I am within the block and that gives me my global coordinate I, right? So, some bookkeeping computation just to get the coordinates here. And then this is important too. You see this pattern basically in all the CUDA code that people write. there's no kind of out of bounds checking naturally. And so what you do is I have my coordinate and I'm going to check to make sure that you know I am supposed to be processing something that's inbounds. And some of the threads at the very end of your block, they're going to be processing stuff that's out of bounds in memory. And you do not want it to touch those. And so you basically condition it on i less than num_elements. And you do nothing if you're outside of that. Sorry. Yes. Sorry. This is just the extension that you sort of write the CUDA code in. It's to distinguish it from, you know, just your standard C code. Okay. so this is just a file name thing is this CU. There's nothing particularly special about it. okay. And then so now you know within here we're going to just do our computation,

56:23 Right? It's just going to be I'm going to write out I have my input in. I'm going to index into the E element and I compute my GeLU just like I did before and I assign it to out of I and then I'm done. Right? That's all that's all that I need to do. And since this is all pointer stuff, I don't really need to worry too much about what is kind of actually happening here. So that's basically it. I can then take my sort of CUDA GeLU code that I have and then I can load this sort of C++ code in line and then I can just have it compile into a module all within Python. It's all very nice and convenient. You don't really have to go out onto the command line and do things. And so now we have CUDA GeLU defined. so this is nice and basically it's a compilation of this. and I can call it from within Python and we'll use the C bindings to call this guy. Okay, we're done calling CUDA GeLU. I have my, you know, I can check that the manual GeLU and the CUDA GeLU are the same. And now let's benchmark the two. so I have the time that it takes to run PyTorch. And, you know, just like last time, it's about 1.1 milliseconds. and manual time, remember, is 8.1 milliseconds. And so, drum roll, what is our CUDA time? Well, we've gotten it down to 1.8, right? Not quite as good as PyTorch's implementation, but, you know, we're we're getting pretty close to PyTorch time, right? We've we've gone from 8 milliseconds to 1.8 milliseconds, which is which is not bad. because that C code wasn't that hard to write. And so now we also do some profiling. and we can kind of see what is happening

58:00 Here now. and you know it's called the GeLU kernel, right? This is the code that got shipped off to the GPU. and then it's calling empty_like this is the initialization. and then empty_strided, right? and then cudaLaunchKernel and cudaDeviceSynchronize. and that's basically all that's happening. And notice how you know once again this is a single CUDA kernel eats up 100% of the GPU time. Kind of like what we what we wanted, right? Okay, so there's some further optimization we can do, but this is really already solved the problem of you know kernel fusion. We fused all the operators together. Okay. so pretty good. these kinds of elementwise operations are easy to write in CUDA. Like if you have a new kind of I don't know nonlinearity. You could easily write a CUDA kernel for it yourself if you really wanted to. but more interesting operations are going to require reading multiple values like doing reductions. Those are going to get a little more complicated. FlashAttention will be a little bit more complicated but not too much so when you have to do it in the assignment. Okay. any questions on the on the simple C++ CUDA kernel? Yes. Check the beginning. Yeah. Does that throw an error? Is it like caller kernel? Yeah. So the question was what happens if it's not contiguous? At least in the code that we wrote it will just throw an error because it's an assert. you could potentially write code to handle it, but there's almost no reason for memory to be fragmented because it will allocate contiguously. and you won't deallocate like the middle of a memory unless you're doing something like really tricky. and so you should you should really unless you're doing something pretty advanced expect to have contiguous memory.

59:44 Sometimes you do like a transpose or jump operation that makes memory not. So like when you're encoding at a higher level should you be careful to conversely make like forced to be continuous before calling operation. Yeah. So the question was like if you're transposing then you're no longer going to be contiguous. You're going to have like a you know jump between all the elements in the index. If you're sort of row traversing something that's sort of column stored. yeah. So I think transpose or like views or like essentially shuffling dimensions is like the one exception to this. But that's handleable in like the outer like sort of the wrapper part, right? You can basically pass it something that is contiguously indexed. and for a lot of the matrices, you won't really care, right? So yes, what would happen if you were to choose a different block size, right? So what would happen if you chose a different block size? the sort of GPU related sort of concerns would kick in. Sort of like do you have enough blocks for to saturate your SMs? and do you have enough work within each block? And those are like kind of the two things that could matter here. But I think my guess is that for block sizes that are relatively large like 1024, it probably won't matter past a certain point because we're not doing anything advanced. It's all elementwise operations for this like very simple example. yeah. is the reason that our non GPU version was so slow because this ask to like do a small operation of GPU back. So, so the question was like why was our non CUDA kernel sort of like manual

61:20 Thing so slow? it's not that it's sending things back from GPU to CPU per se like X is going to live in the GPU. we allocate it in GPU like we'll do like as the device like CUDA but it's going to basically not be in the SM the whole time right so once we do like X squar right that's a you know a CUDA kernel and so that multiplication operation will read the sort of vector from the global memory into the SMs do the computation it'll write it back and so this is all in the in the sort of DRAM to SM communication cost

Triton: the same kernel, in Python, in blocks

61:53 Rather than the CPU to GPU communication cost of course, if you write like as device CPU, then you'll hit get the you know CPU transfer cost in addition to the to the DRAM transfer cost. Okay, so now you've seen that and like okay so that was not too painful but it would be really nice if we had nicer sort of Python abstractions for writing CUDA kernels and this is what Triton is and Triton is quite nice. It like has this very nice middle ground where you don't have to manage literally everything about the GPU. So Triton is sort of a domain specific language developed by OpenAI in 2021 and it makes GPU programming much more accessible. So like you write everything kind of in Python and you don't really think about the threads anymore. You think about thread blocks and Triton manages a lot of stuff that is annoying but can be automatically optimized. So it can manage coalescing of memory. so remember that you know from VRAM you get four sort of adjacent values at once with something called burst mode. So you really want to make sure that you know your memory retrievalss are sort of grouped into adjacent sort of four element or more sort of calls at once. So it will handle those automatically. It will group those. it will do shared memory management when you need to sort of manage which sort of memory that you're writing to within the SM with multiple threads from within each SM you know you might need to stop or start threads all managed automatically but scheduling across SMs or what different SM do

63:36 That's manual so like the kind of the programming model is that you're going to think kind of at the SMcententric level and the compiler will handle a lot more of the lower level details and Triton is quite nice because it can outperform by quite a bit a lot of PyTorch implementations. So, it's kind of like going all the way to writing CUDA, but you're still in the very familiar Python land. And I think a very underappreciated advantage is sort of as it's written here. It's all in Python. You can step through it. You can kind of debug it fairly nicely. And so, let's step through a Triton kernel. Like once again we're going to write GeLU and we're going to do it in Triton. So this I've you know put the code to be as similar structure as possible to our other code. Right? So this is sort of the CPU side code so to speak. This is the wrapper code. It takes in X which is a torch tensor and I've got my two asserts at the top. and I'm going to allocate an output tensor Y using empty_like once again. And it has the same exact sort of coordinate computation. sort of components and even the kernel launch looks very similar. I've got this num blocks annotation and then my block size is you know at the end here not in part of this brackets but basically I'm passing the same information to my kernel and now trying kernel is this code over here and this is going to do the same thing as what we were doing before but now it's nicely written in Python and you know the mental model here is the inputs are going to be at x_ptr y_ptr is the output vector sort of the starting

65:14 Coordinate and the block size is how big you know each of my blocks are and num_elements is going to be sort of the very end of my array. So now I need to get this set of lines 557 to 561. This is doing the computation of my index right I did I equals you know some formula before this is doing the same calculation over here. I'm calculating where is the start of my current block. Well that's my block ID times the size of the block. that gets me. Let's say I live in block one. It'll get me this point right here at the middle. and then afterwards I need to know where do I live within my block? Well, that's going to be kind of the offset. But now notice one difference. I don't get in an offset because I'm not programming threads, right? I'm programming blocks. And so what does that mean? Well, my offsets are actually a vector, not a single value. because this is basically going to be I'm going to do vectorized operation where the vectorized operation is going to be handled by different threads. So here my offsets are the start of the block plus a vector this range of block size sort of offsets. So I'm my offsets are all of these coordinates within block one at once. Of course, if I'm at the very end, I might go off the edge. And so, I need a mask to handle anything that lives off the boundary of my vector. Now, I'm going to load in a sort of single vectorized operation everything at once. So, x_ptr plus offsets. These are sort of the values that I'm responsible for masked up and it's loaded into X which is my sort of internal values my internal sort of

66:56 Temporary vector that I need and with this temporary vector I'm going to do exactly the old GeLU computation. there's no tanh so I compute that manually but this formula you can convince yourself is the same as what we have here. and then y is going to be the formula computed up here. Now once I'm done I need to write it back into my output sort of buffer or my output vector and so I compute sort of my targets. So this is y_ptr plus offsets. I take my values my temporary values y and then I store it right. So this is very similar to what came before but this one is the vectorized version. I get to operate on an entire block at once. And so instead of kind of thinking at the perspective of a thread, I'm thinking from the perspective of a block, but not too different, right? This is all fairly similar stuff. So now I've written my Triton GeLU and all right, I will I will do this fairly quickly. All right, so one last thing I will only point out a few things here because I don't want to get like so in the weeds that you all like get up and leave. but the one

Reading the PTX that Triton emits

68:02 Last cool thing that we can do is Triton of course compiles into low-level sort of almost machine code for the GPU. And we can look at, you know, this very low-level called PTX code after the Triton compiler sort of goes over it. And it's actually kind of cool. You can kind of see how the GPU like actually works at the thread level. So this is the Triton GeLU kernel. It was generated by the compiler. And at first it's going to do some of the really basic stuff. So what's it doing here? It's saying, well, I'm going to need to store some values, right? I'm going to need to store intermediate computations. B means actually sort of untyped sort of basically like bytes. So I need bytes that are sort of 32bit size. I need floats for doing computations called f. And I need another set of registers that are 64 bits. And you know that's another set of registers. and so I have all these sort of registers that I need for temporary computations. And then starting here I'm going to start computing basically my coordinates. So sorry this part is loading the the various arguments to the function. So things like the x_ptr and the y_ptr get loaded here. I starting here I start computing the coordinate offsets of my Triton sort of kernel. And then once I get down here, this ld.global, this is the code that's used to load the values from x_ptr back into my temporary registers. So it's basically saying load %r2, %r3, %r4, %r5 using the memory position in %rd1. And notice how it's loading four

69:44 Things at once because it's cleverly handling coalescing, right? We know we can get four values for free. we should, you know, operate on all four of these values at once because we get them. And then you do the same thing again for you do the same thing again here. And then you start to get basically the floating point operations mul.f32 which basically goes through and does the tanh computations. I'm not going to explain all the different pieces, but you know here it's doing it's multiplying by a constant. It does a x to the cube like multiplying the same numbers multiple times. and then it's going to compute here, you know, 2 to the x, but we want e to the x. And so it multiplies by log two to get the exponentiated base. You can really see all of the different like literal step-by-step operations that the GPU does in order to get you the final result. And so I'll skip all over to the end. This is all floatingoint computations that it needs to do. And then at the very end it stores the values that it has %f38 through %f41 into %rd4 which is the memory position of our output. Right? So this is kind of like what's actually happening at the low level. and we see that each thread is operating on four values at a time and its temporary storage is the registers which is the really high-speed storage that it has very locally. So we can see you know this is going to you know just looking at it be probably pretty fast code right. Okay. So that was the PTX and we can you know go through and see what it's doing for all sorts of things. But now let's go back and actually benchmark things. So we got manual GeLU 8.1 seconds, PyTorch

torch.compile, and when you can still beat it

71:22 Time 1.1 seconds, CUDA time 1.84 seconds, Triton time 1.848 seconds. So we didn't get any faster, but it was much easier to write Triton code, right? We wrote it in Python. We thought about blocks. We could do vectorized additions. if you're doing more sophisticated stuff, you know, it basically Triton will handle a lot of the memory stuff for you. and so it's actually pretty good. And then profiling once again, we see single kernel launch that consumes all of the GPU time, right? So that's great. and that gets, you know, Triton kernels. The last thing, at least in this sort of Whoops. One second here. Okay. that I want to talk about is torch compile. of course writing CUDA kernels is cool and it makes you feel really good. but maybe we don't need to do that, right? Like the things that we were doing here were very simple. We were just taking these like you know x cubed and like exponentiation operations and we were just shoving them all into a single CUDA kernel. And so maybe we can just do that without you know doing much. And so, you know, we've had the several different ways that we've showed you, but the last one I want to talk about is this thing called torch compile, which will take you know, nonoptimized PyTorch code, and it will write more optimized code. And so here it's going to attempt to automatically do optimizations like kernel fusion. and this compiled GeLU is going to be, you know, equivalent in the actual outputs that it generates. But now let's let's look at the run times, right? so we've got some runtime variation, but basically

73:04 The same kind of numbers, right? 8.1 seconds manual, 1.1 seconds PyTorch, 1.8 seconds, and then 1.47 seconds on torch compile, right? So the punch line here is modern JIT compilers are pretty good. It can do optimizations like operation fusion without you having to do very much at all. And if you look under the hood, you can kind of see that there's basically once again one thing that happens. This is a sort of fused add multiply tanh Triton code. So it's generating Triton under the hood that basically is doing similar kinds of things as our Triton code, but it's actually slightly more optimized than what we did. And so it's getting slightly better performance than even our code. So torch compile is quite nice. Yes. How do you feel like compiled? Like you're going to like try to implement your price version like it can't do flash in right. Yeah. So the question was like when do you know that I guess maybe the better way to phrase that question is when do you know you can do better than torch compile right is sort of the relevant question. and I think for simple stuff like simple operator fusion or the other thing that it's very good at is optimizing matrix multiplies. so torch compile as I said before can do things like if it knows the shape of the matrices can figure out which kernels to dispatch. It is very good at those things. I doubt that you can get much better than that. But there are things like if you've seen FlashAttention-1, 2 and 3 those are pretty non-trivial

74:45 Optimizations like these days torch compile and like JAX's XLA compiler can do those but that's because we know in hindsight that those are the right optimizations to do. I think some of those things are a little bit non-trivial to figure out like FlashAttention-3 has additional sort of hardware level optimizations that leverage you know the H100 hardware that's not obvious to do with a JIT compiler. and so there are some things that I think are quite hard with torch compile that I think you could do better. But in general, like I think the point here is, you know, you

Fused softmax: the first kernel with a reduction

75:16 Shouldn't go home and say, I'm going to CUDA kernel like I'm going to write CUDA kernels for every single part of my language model, you know, that's probably not a good use of your time. But if you're writing a new architecture with some complicated piece and you're not getting utilization, but you think you can, that's maybe the time to really bust out the Triton. Okay, so we're we're basically at time but we can quickly go through one last example of Triton. Maybe this will be useful for you in assignment two of doing softmax. So one difference is until now we were doing just basic element wise operations and that's really easy because you just operate on each element and there's sort of no sort of complexity to those kinds of things. So now let's do softmax which is it has a reduction operation where you have to add across all the elements. So how do we do that? Well what we want to do is we want to normalize across each row of the matrix and you know what we would like to do is we'd like to make this fast. So a naive version of this is going to be pretty slow. And now we're going to write the Triton kernel. So if I wanted to be lazy, the easiest way to do this is okay, actually you can think for a moment about what the easiest way to do this. Now let's say you want to write a softmax. So you're going to normalize each row of a matrix and imagine these matrices are pretty small. So you're just writing a kernel for small matrices, right? So if you're doing this, what's the right kind of block design? Well, maybe what we should do is our grid should actually just be rows. So each SM is going to handle a single row. That's kind of the optimal

76:53 Thing to do because if we can fit a whole row into an SM, then we just sum across that row in the SM and then we divide, right? That's that's great. And so that's going to be the simple design for our very, you know, naive softmax kernel here. So all we're going to do is that we're going to make the block size basically sorry, we're going to make each block a row. And so the block size should be number of columns plus, you know, a little bit of buffer to sort of be able to fit all the columns. So this is triton.next_power_of_2 of n. And that's a nice way of padding out your columns. And then I'm going to make each block a row. So the number of blocks is exactly the number of rows. And then I have my Triton softmax kernel which is written in kind of the way that you expect. So now we have a matrix rather than a vector. So we have x_ptrs, we have y_ptrs, we need the strides of the matrices. and then we can basically figure out what row index I'm in. I can get the column offsets. This is going to be the same kind of code as before. In fact, getting the row offsets simpler because each row is a block. And then now I'm going to do basically the same kind of stuff. I'm going to load in each row into my sort of SM's sort of local memory. And then I'm going to do computation exactly in a way that looks like a softmax. I have my row. I subtract my max. I take the exponent. I sum it and then I divide which is going to give me my softmax normalized row and I write it back to global memory. Right? No complexity at all. whenever your computations fit nicely in SM, writing Triton code looks very similar to writing just normal Python code just with a little bit of load and store and keeping track of where the blocks are.

78:35 Right? So life is pretty simple. Let's go back. oh wait, where were we? To the Triton. Here we go. And then we can kind of see how fast all of our different pieces of code are. So I'll zoom out again just make sure. Okay, so manual time takes 3.7 seconds. our compile time is 1.3 seconds for torch compile. the PyTorch time is 1.5 seconds. and the Triton time is 1.9 seconds. It's a still a little bit slow. torch compile can actually do better than sort of the native PyTorch implementation especially when it knows about the shapes and sizes of certain operations. so finally we can look in the profiler the manual softmax is kind of a disaster here. You see all sorts of crazy operations happening all over the place. Let me let me clear this if we go back up here. Okay. Yep. we see all sorts of operations happening. you know, we have x, we have max, we have sum because we've implemented things naively and we've got memory reads and writes everywhere. the compiled softmax is just going to be sort of one fused softmax operation that goes quite fast. and then we've got pytorch softmax which is also one CUDA kernel call and same thing with our Triton softmax. We have our nice Triton softmax kernel that is a single fused kernel for everything. Okay, I won't go through the PTX code for this. I think, you know, we're we're kind of at time and I don't want to drag you through that low level again. but hopefully this has given you a flavor of lower level GPU programming for the purpose of making language models go fast. And hopefully you'll

80:13 Have fun doing assignment two. Thanks.