6 ms·
Show HN: Flash Attention in ~100 lines of CUDA
- treesciencebot 3y agoPretty neat implementation. In general, for these sort of exercises (and even if the intention is to go to prod with custom kernels) I lean towards Triton to write the kernels themselves. It is much more easier to integrate to the tool chain, and allows a level of abstraction that doesn't affect performance even a little bit while providing useful constructs.
- ixaxaar 3y agoYou mean triton the inference server or triton the DSL for cuda?
- whimsicalism 3y agothey mean the dsl (not just necessarily for cuda)
- p1esk 3y agoThe DSL: https://openai.com/research/triton https://openai.com/research/triton
- treesciencebot 3y agotriton the DSL.
- whimsicalism 3y agoyeah even the official flashattention is moving many implementations from cutlass to triton except for the main mha backward/forward pass
- jart 3y agoIt was written with cutlass? No wonder Peter Kim found it valuable and worthwhile to de-obfuscate. Adopting a new programming language invented by OpenAI doesn't sound like a much better alternative. I'd be shocked if either of them were able to build code for AMD GPUs, where it's easy to adapt CUDA code, but not if it's buried in tens of thousands of lines of frameworks. I like open source code to have clarity so I can optimize it for my own production environment myself. When people distribute code they've productionized for themselves, it squeezes out all the alpha and informational value. Just because something's open source doesn't mean it's open source. I think people mostly do it to lick the cookie without giving much away.
- synquid 3y agoTriton has an AMD backend, although work is still ongoing.
- imtringued 3y agoYou will also be able to use Triton to target Ryzen AI.
- fpgamlirfanboy 3y ago> allows a level of abstraction that doesn't affect performance even a little bit The second part of this sentence is true because the first part is false.
- treesciencebot 3y agozero cost abstractions exist. doesn't mean all abstractions are zero-cost. or being zero-cost somehow invalidates their abstractness/genericness. but maybe we differ on the definition of abstractions.
- fpgamlirfanboy 3y ago> zero cost abstractions exist So does perpetual motion :shrug: but my point is Triton is not an abstraction in the least. Source: 1) I spent 6 months investigating targeting other backends 2) Phil himself said he doesn't care to support other backends https://github.com/openai/triton/pull/1797#issuecomment-1730112311 https://github.com/openai/triton/pull/1797#issuecomment-1730...
- deleted 3y ago[deleted]
- fpgamlirfanboy 3y agoIt's amazing how heavily provided hn is. I have a response here that's been deleted that is like 15 words, including a link to source that corroborates my claim but that response contains a transcribed emoji and so it's been deleted by dang or whomever. Lol super rich environment for discourse we've got going here.
- queuebert 3y agoAs a person who finds CUDA extremely easy to write and integrate, what does Triton have to offer?
- whimsicalism 3y agoblock level rather than thread level programming, automatic optimization across hyperparameters, makes it much easier to write fast kernels
- araes 3y agoFor those who have no idea what's being discussed, quick background. Discussing: Transformer [1] memory issues and approximate attention [2] in machine learning training. Specifically: FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. [3] As a side comment, this entire industry is sorely in need of at least intros. The entire space has moved so fast in the last year I need an entire new dictionary and thesaurus for all the terms they've created. Notably, because of this, found out Google has a glossary of machine learning terms. Actually somewhat handy. [1] Google Machine Learning Glossary (Transformer): https://developers.google.com/machine-learning/glossary/#transformer https://developers.google.com/machine-learning/glossary/#tra... [2] Same (Attention): https://developers.google.com/machine-learning/glossary/#attention https://developers.google.com/machine-learning/glossary/#att... [2] arXiv: https://arxiv.org/abs/2205.14135 https://arxiv.org/abs/2205.14135
- deleted 3y ago[deleted]
- robrenaud 3y agoRegarding your comment about how fast the research and industry is moving, would HN readers be interested in relevant one or two paragraph summaries that are basically "explain it like I am a machine learning engineer from 2020" but also knows the power of these models from a perspective of using ChatGPT or MS Copilot? That is, assume a fair amount of technical knowledge about the fundamentals, but don't assume that the reader is paying any attention to have whitebox knowledge of the current state of the art.
- jprete 3y agoI personally have been looking for "explain it like I'm a CS PhD with lots of experience and the ability to look stuff up". But I suspect your summary would be pretty handy as well.
- jhanoncomm 3y agoI reckon you need tacit knowledge. Experience. Luckily in the order of 100 hours not 10000. Build a GPT using Python and Pytorch. For a good course: Andrej Karpathy is your keyword. At $1000 his course is great value. But actually it is free which is even better ;-) It wont take you to flash attention but will ramp you to the point you could probably read papers about it. I almost got that far then life lifed me. But I was able to implement changes to the architecture of GPT and do some “hey mum I am doing SOTA (2021) machine learning”.
- saiojd 3y agoWhat does __syncthreads() do here exactly? I'm new to CUDA, could get the overall idea of the FlashAttention paper but not the details.
- cavisne 3y agoCauses every thread in the block to wait until they have reached this point. Worth reading a cuda primer for more details on blocks/warps. Since the threads are relying on each other to fill the SRAM with all needed data if you didn’t wait then values would be missing.
- xrd 3y agoAny CUDA primer you recommend in particular? I had this same question.
- winwang 3y agoHere's an article on syncing in CUDA via cooperative groups: https://developer.nvidia.com/blog/cooperative-groups/ https://developer.nvidia.com/blog/cooperative-groups/ There's also explicit warp synchronization, i.e. __syncwarp(). More on warp primitives here: https://developer.nvidia.com/blog/using-cuda-warp-level-primitives/ https://developer.nvidia.com/blog/using-cuda-warp-level-prim...
- cavisne 3y agoProbably https://www.youtube.com/watch?v=nOxKexn3iBo https://www.youtube.com/watch?v=nOxKexn3iBo (or just skimming the attached colab).
- xrd 3y agoThis is terrific, thanks!
- Linda231 3y ago[dead]
- zer0zzz 3y agoThis is fantastic. I am just starting in the ML space (compile from compilers) and I love short kernels that I can use to understand things better with.
- lagrange77 3y ago> compile from compilers What does that mean?
- zer0zzz 3y agoTypo, meant to write “coming from compilers”
- einpoklum 3y agoMy GPU work is not in ML (deep or otherwise); but ... 1. "100 lines of CUDA" + PyTorch; maybe this is useful and maybe it isn't, but counting lines of code on top of a huge codebase is not very meaningful. 2. Launching separate kernels, synchronously, on the default stream, for various operations, is typically not the right way to utilize a GPU.
- chillee 3y ago> maybe this is useful and maybe it isn't, but counting lines of code on top of a huge codebase is not very meaningful. In this case it's pretty reasonable imo, since the kernel itself is fairly independent - the usage of torch is just for some bindings for the data structures. > Launching separate kernels, synchronously, on the default stream, for various operations, is typically not the right way to utilize a GPU. This is actually the standard way to do things in ML. Assuming you're from a HPC background (where this may seem quite strange), the biggest change is that "More or less everything in ML runs on the GPU", so there is very rarely any device to host synchronizations. In addition, each individual kernel is typically run on fairly large chunks of data (a million elements would be on the smaller side), so maximizing occupancy with streams is not as necessary as in HPC.
- danielhanchen 3y agoFantastic work! Extremely neat and clear implementation! Interesting note on the backward pass - what do you think are the main blockers for a backward pass?
- tspeterkim 3y agoThanks Daniel. The main blocker is me not able to fully grasp the backward pass. (trying to understand Appendix B.2 in the original paper) I need to get more comfortable with matrix derivatives before I can confidently reimplement it in the same minimal way as I did with the forward pass.
- danielhanchen 3y agoOh ok! Ye the backwards passes are always much more difficult due to the derivatives!
- dcanelhas 3y agoIf CPU/GPU execution speed is the goal while simultaneously code golfing the source size, https://halide-lang.org/ https://halide-lang.org/ might have come in handy.
- jobsrajasthan 3y ago[dead]