Weight Folding, CUDA Streams, and the Bug That Made My Model Speak Backwards — Filip Makraduli

summarized

TLDR

A paper co-authored by Filip Makraduli proposes two algebraic tricks—weight folding and deferred normalization—that speed up RMS norm in transformers by fusing normalization into matrix multiplication and parallelizing operations, reducing wall time despite minimal math contribution. The speaker also describes a CUDA streams bug that caused a one-step generation lag due to implicit stream join, fixed by explicit synchronization. The techniques are implemented in the Transformer Tricks repo and can be deployed via Superlinked's inference engine.

Key points

The paper introduces weight folding and deferred normalization to reduce RMS norm wall time in transformers.

RMS norm accounts for significant wall time due to GPU overhead from starting operations multiple times.

A CUDA streams bug caused a one-step lag in generation because of an implicit stream join leading to race conditions.

The fix required explicit synchronization of CUDA streams to ensure correct buffer reads.

The techniques are compatible with torch compile, flash attention, and quantized models, and are available in the Transformer Tricks repo.

Techniques

  • weight folding
  • deferred normalization
  • weightless normalization
  • canceling pre-normalization
  • CUDA streams parallelism
  • explicit stream synchronization
Transcript (captions)

0:12 Hello everyone. Uh Thank you for coming and uh I'll start the talk now. So This talk is uh around a paper

0:27 that I did um which is very simple. The proposition is very clear. It's basically

0:37 two lines of algebra that make uh the RMS norm layer in transformers cheaper, quicker, and kind of improve improve it as like a layer in the transformer architecture.

0:53 Similar to how layer norm once used to be the standard and then it was substituted by RMS norm. This follows along um this way of thinking.

1:06 And I got the chance to kind of meet some people from the open source world and um I co-authored this paper together with uh Niels Graf who was the kind of the creator of this.

1:20 And the work follows from there. So this is uh presented on archive. You can have a look, read it, test it out. There is a repo as well. And the concept the

1:35 let's say the idea and the way of thinking it's easiest to explain with maybe flash attention. So in a similar way of how um flash attention kind of weights until there's a multiplication

1:49 and tries to limit this uh communications between memory so that the whole process is faster. This is kind of a similar thought along those lines and it does certain improvements

2:01 that make the RMS norm process much quicker and in in effect improve the whole

2:14 transformer. And one question is okay, why RMS norm since that layer does almost none of the math? And that's true. So the share of the

2:27 kind of math portion if you look at it is quite small. However, the clock time or wall time as they say is quite big and

2:38 for example in one decode step. So right when like inference is performed the RMS norm can be started like 33 times. Of course, it depends on the model and so on. In the paper, you have

2:51 the specific models and how this was tested. Um and the question is how this can be improved and how

3:00 this wait for the matrix multiplication can be kind of avoided. And the reason why this is slow is because the GPUs are not slow or bad at math, but they're bad at

3:16 everything else around the actual math. So that means starting the work, the actual work. So for example, starting the process as it happens in some of the experiments

3:31 33 times that takes a long time and for example, fusing um each normalization into the matrix multiplication can help avoid this. Also

3:44 doing weight folding can help in kind of moving data between memory and um that's a process that's also slow for GPUs.

3:55 And also waiting. Uh so, for example, deferring the division that's done in the RMS norm layer is also a way to avoid this waiting step. So, basically, what this paper does is it improves all

4:11 these three aspects by doing a few algebraic tricks in the way RMS norm is computed. That's it. And math-wise, these are the tricks.

4:25 Uh it's mainly around the first two propositions. One is weightless normalization. Uh you can see that here. Um and deferred normalization, so um that's the second one. And now in more

4:38 newer architectures, there is a situation where um RMS can be kind of can appear twice. Uh for example, in Gemma 4, this happens. So, canceling the pre-normalization also

4:52 works. Um and all of this is algebraically proven in the paper. And the first proposition is this where kind of the the gain and the weight fold folded to one matrix W, uh you can see

5:07 here with an asterisk. And that is computed offline, similar to how maybe in flash attention, you compute some stuff on the side so that there is no uh communication between

5:19 memory all the time. So, this is one step that's kind of um done, this weight folding. And the other step is um deferring um the scalar the scalar divide of the

5:32 matmul so that they can be done in parallel. So, in a normal case, you would have to compute once, then wait, and compute again. In this case, the idea is to kind of split this so that it

5:44 can be parallelized. And the third one, which is kind of a version of this is that um there is kind of if there are two

5:56 um because this is scale invariant, one of them can be dropped and this still works. And this is applicable to newer models um that can have this architecture and

6:07 implementation. So, in order to make this happen in real life, especially this proposition number two, um

6:18 so for this one, for example, it's easy. There is a repo called Transformer Tricks. You can just apply this to any model and it works. But in order to do this, there is some

6:28 kernel work. So, it's not as straightforward to do. So, in order for me to do that, I was implementing this and I

6:40 came out with this experiment once. So, it looks okay in general, where it's like, "Okay, the prompt is the Transformer architecture revolutionally revolutionized NLP because and then

6:51 there is some kind of expected output." But in the output I got, I saw this repetition and one-step lag, as you can see here, the word because appears again. And

7:03 there was something happening with the GPU streams and I was trying to figure out what was happening. And I was getting this one-step lag and kind of um outputs that were from the

7:16 past in a way. Um and in debugging all of this, I realized that um in the process of building something like this, so

7:27 as I explained the proposition two or deferring these two operations, um in CUDA, you can do two things. You can do like tensor cores that do one part of the matrix multiplication and

7:39 you can do CUDA cores that kind of run stuff like element-wise operations, reductions, square roots, and so on. So, the idea was to do this in parallel and get the benefit

7:52 of what I was explaining in the paper to actually test out this concept. So, this is how it was supposed to look like. So, there is if you do things sequentially, there is this idle waiting

8:04 time when you when the vector unit computes the RMS and scaling, and then there is a matrix multiplication. So, the idea was okay, with flash norm, which is the technique in the paper,

8:16 you're supposed to do those both in parallel. So, the matrix unit computes the matmul and the vector unit computes the RMS. So, in that way you save uh time. However, you cannot just do this

8:28 in Python, you have to go a bit lower. And I did that with CUDA code like this. And this looked in general okay at my uh at that time.

8:42 However, um I realized that I did something slightly wrong. And that thing was that

8:50 the join in the end, where you're supposed to join the two streams, was implicit in my case. And when I tested this out, the unit tests worked, the quality seemed

9:02 similar, like perplexity testing, and so on, because it's just like um similar generation, but over long generation, I was able to see this problem. So, I had no idea what this was.

9:15 And the reason was that when I was doing this uh implicit um join, basically, one of the streams hadn't finished the work, so I got race

9:26 conditions that kind of read the past from the unfinished matrix multiplication. So, the idea that I had to fix this was um around the fact that I had to be

9:38 explicit about the join and wait until one of the operations is finished so that I'm certain that when I join I'm not reading from the past. So, that was the realization um in this

9:51 exploration of CUDA streams. And this is how I had things done. So, the join was implicit. So, the post um scale read like an old uh buffer value. And how this is fixed is with this where

10:09 basically you need to mark the end of the matrix multiplication, then mark the end of the RMS, and then post scale wait for

10:20 the first stream and then wait for the second stream. And that fixed the bug and made kind of the paper work and the model speak forwards instead of backwards.

10:33 And that was the cool maybe academic perspective, but I also wanted to try things, right? Deploy this, test it out, see how I can make it work um

10:45 in maybe a more production setting. And you can also read the paper and see all the tests. Um some are done most are done around llama models, but like this works for other

10:55 architectures as well. Um so, what you can do for this specific paper is um for example, the weight folding that I explained the pre-position one, you can just do it

11:07 with um some code in the repo that's like flash you say flashify and it does that. However, with this second thing that I mentioned, you need to do a

11:17 bit of kernel work if you want to do that uh like I explained in my example. And these are some results that are based on llama models and there are different kind of details that you can

11:27 have a look at as well as well. Like what happens if you do only the third normalization, what happens if you do a full fused kernel. Um so there are a lot of experiments of going lower here to

11:39 test all the propositions, and this have been our results um in different, let's say, levels of um scrutiny and detail. But even the simple one with like weight

11:50 folding um shows some improvement. And this also works with like the day-to-day tools that you use in the models. It's not like you have to

12:00 reinvent the wheel or, you know, do things from scratch. So it works with uh torch compile um because the it's kind of like a new checkpoint, and that's it. Flash

12:13 attention does similar tricks at a different layer, and also it works with quantized models. So it's totally cool to actually apply this, and you can get a model that has this cool new

12:25 normalization layer. And where you can get this um details and code to actually run this is this transformer tricks repo. So

12:35 uh it has different algebraic tricks like I explained, as well as this paper that I mentioned. And also there is the GitHub uh not the GitHub, but the Hugging Face uh model

12:47 repo where I've done this with some models, and you can have a Hugging Face link to that model and test it out. Um and what you also can do with this Hugging Face models is to deploy them in

12:59 production. And so when I was thinking about doing this, um I realized that, okay, now that, let's say, the science is done and there is a link to a Hugging Face model,

13:11 um Superlinked's uh inference engine was a cool way to actually deploy any um uh Hugging Face model, and we've done this at hackathons where people would bring like a custom Hugging Face model

13:24 or checkpoint that they have with their fine-tuned stuff, and you can test out like even if you have some version of this algebraic tricks that you want to improve a model and test on test your

13:36 own research ideas, you can actually try that out and have a deployed version on a cluster of this model and not have to worry about this glue code around deploying models.

13:49 So, that's um pretty cool. And the the point is that if you have the full cluster open source and the model inference open source, you can actually test out this kind of

14:00 maybe more novel research ideas where if you want to do kernel manipulation or

14:10 flash norm and things like that, it's much more difficult to do that do this at a rented endpoint where you don't own the inference. It's you want something that's portable and flexible to actually

14:22 allow you to do this stuff, but it's also production ready enough so that you can test things out at scale. And you can, for example, use site to combine this with other models like, as

14:34 you can see in the top left, there is you can have this flashified models with different other models to do agentic tasks if you want and kind of do that end-to-end bigger use case.

14:48 And the way site works is this production cluster helps you deploy the models, so you can have a look at site's repo as well for more details on this. Um and also there is a smarter queuing

15:00 mechanism that helps you, especially if you work with smaller models cuz when doing the flash norm stuff, I worked with like smaller llama models and also with small agents from hugging face. So,

15:12 having a way to deploy smaller models that can also work on like the same GPU so that you don't have to spend your money on GPU cost, but

15:24 actually kind of switch models around, especially smaller models. It was quite useful. And you can also control the model configs through an API as well as the

15:34 cluster, which is also pretty convenient without having like an infra guy supporting you in your open source research. So, that's cool as well. Um, and you own your cloud, which is

15:46 useful if you want open weights, open models, open source. And there's also like a catalog that Sci has of different models, um, not just the ones I mentioned, but you can have a

15:58 look. There's also re-ranking embedding models if you're building something along those lines. And with that I'm kind of finishing this story of my research journey where

16:10 I co-authored this paper, um, around the technique that improves the transformer, but also found a way kind of to bring this to, let's say, production and test it out and find a

16:21 way to play around with this open source models. And feel free to contact me on LinkedIn, maybe if you have any questions or contributions. A lot of this stuff that I've mentioned, like

16:32 some of them are PRs on like vLLM or on Hugging Face. You might find them all around. You can also see the check out the paper. That's the archive link that you have there. Um, and you also have

16:44 the Sci repo and my LinkedIn. Um, so thank you very much for attending. >> [applause] [cheering]

16:57 >> And you can catch me for questions. We'll be here, close by. >> [music]

Frontier News · by Hyperjump Technology