AI’s compute fragmentation: what matrix multiplication teaches us
modular.com
modular.com
https://opensource.googleblog.com/2023/03/openxla-is-ready-t...
> OpenXLA is an open source ML compiler ecosystem co-developed by AI/ML industry leaders including Alibaba, Amazon Web Services, AMD, Apple, Arm, Cerebras, Google, Graphcore, Hugging Face, Intel, Meta, and NVIDIA. It enables developers to compile and optimize models from all leading ML frameworks for efficient training and serving on a wide variety of hardware
I used to think this. And I think, in theory, it is true. But the fact of the matter is, modern ML just doesn't use that many kernels. Every framework uses the same libraries (BLAS) and every library uses the same basic idea (maximally saturate FMA-like units).
Large language models are being run natively on commodity hardware with code written from scratch within days of their release (e.g. llama.cpp).
From a conceptual standpoint, it's really easy to saturate hardware in this domain. It's been pretty easy since 2014 when convolutions were interpreted as matrix multiplications. Sure, the actual implementations can be tricky, but a single engineer (trained in it) can get that done for a specific hardware in a couple months.
Of course, the interesting problem is how to generalize kernel generation. I spent years working with folks trying to do just that. But, in retrospect, the actual value add from a system that does all this for you is quite low. It's a realization I've been struggling to accept :'(
All of the convolutions we have run on kernels, some pre-built and customized/chosen from a list based on performance, and some dynamically generated. PyTorch 2.0 for example decomposes and fuses operations, then uses OpenAI's Triton to dynamically generate a custom fused kernel that tends to be very efficient.
There are still hand-written kernels, even, Flash-Attention and the Memory-Efficient attention papers both caused huge leaps forward because they manually went through a lot of the inefficiencies of naive matrix multiplies for attention w.r.t. the hardware design and optimized it quite a lot.
I think generalized kernel generation though may have more life in it than you might suspect! It is a fascinating field and I do not know nearly enough about it. I hope someday to be able to write my own Triton kernels/get to know how it integrates as a dynamic compiler for PyTorch code. We certainly live in wild times. Crazy indeed.
Triton is used in a templated way for a very specific albeit pervasive hardware (PTX compatible GPUs), which is why it works so well. Here's some of the code: https://github.com/pytorch/pytorch/blob/a66625da3bcdf1e262dd...
Generalized kernel generation (i.e. synthesis of optimal performance from non-expert user defined kernels and novel hardware) would be fantastic to have, but it just doesn't seem particularly necessary in the field.
> Sure, the actual implementations can be tricky, but a single engineer (trained in it) can get that done for a specific hardware in a couple months.
I want to agree with you on this, but in practice, it's...
1. Hard to hire that engineer with deep expertise in handwritten kernels. CUDA engineers are still hard to come by and doesn't scale with productionized AI engineering demand.
2. "A few months" is a tough pill to swallow from an engineering roadmap POV, especially when models are deployed on a monthly basis. Most of the hand tuning efforts aren't scalable and will have to be done again on most iterations. This is especially true in reinforcement learning and robotics.
> But, in retrospect, the actual value add from a system that does all this for you is quite low. It's a realization I've been struggling to accept.
Yeah I remain neutral on this. On one hand, I can see that especially having to invest significant engineering effort (see point 2 above). On the other hand, you won't really know until you start benchmarking these models (and as you should).
By committing it to a common library that a lot of people use? There are already multiple libraries with optimized matrix multiplication.
This is also exaggerating the expertise required. I'm not going to claim it's trivial, but you can genuinely google "intel avx-512 matrix multiplication", and find both papers and Intel samples.
Naively, I wonder if this is the kind of problem that AI itself can solve, which is a rather singularity-approaching concept. Maybe there's too much logic involved and not enough training data on different configurations for that to work? A bit spooky however, the thought of self-bootstrapping AI.
Edit: I mean when you still see papers every year with large improvements in perf, and things like 'we used tensor cores and managed to get back fp32 accuracy with 3 rounds of the things' - what? - I can attest it doesn't take 2 weeks to get this kind of results. And it's just getting started on tensor cores! And when on the nvidia forums someone says 'nah probably no improvement to use tensor cores for fft' and you get a link with a paper with a significative improvement in perf using tensor cores, I say we're just starting.
Speaking of GEMM fusion that you mentioned, flash attention is basically GEMM fusion with online softmax right? This is something I believe really cool and can be made really easy wit a proper abstraction. Say, you may move a chunk of computation under a certain loop and instruct the compiler to optimize data movement or cache intermediate tiles somewhere on chip
Cutlass is supposed to be the first step and to anyone who struggles to understand WTF you're doing when using it, you are not alone. I've seen literally amazing room-silencing stuff with it, but heavy template stuff is really not my thing.
Just like you said, really appreciate that we could actually understand what is going on internally inside the kernel with cutlass, and customize it in a way that cuBLAS doesn't necessarily provide.
Have to agree with you that the template stuff is really annoying. Well, even with some template tricks, the error messages are still less readable, and it is where I think a better abstraction could benefit. Imagine you have an abstraction that simplifies the those threadblocks/warps/etc in a unified way while generalizing it to more backends (AMDGPU, Vulkan, AVX512-VNNI, etc), providing more friendly error messages along compilation given the abstraction is almost certainly more structured than pure c++ code.
I love MLIR and Modular, so please do share more about it! If it’s potential distraction from this thread, I’m also open to email communication if you are interested!
Oh btw, to clarify, I’m not saying Triton is an ideal abstraction. I love it and it’s super popular because it’s the most user-friendly option for ML researchers to write performant kernels on certain gpus, but from a MLSys researcher’s perspective, I’m personally more ambitious and wanted to target broader range of hardwares. Also I really appreciate Philippe’s work that makes Triton really performant and easy to use.
Hey are you referring to 3xTF32 (https://github.com/NVIDIA/cutlass/tree/master/examples/28_am...)? IMO this is a perfect example where proper abstraction could save engineers non-trivial amount of time - imagine a compiler stack which allows 3xTF32 as a normal dtype and subsequent analysis compatible with this special dtype :-)
man this is such a funny closing comment - what exactly do you think is involved in designing a compiler that enables devs to optimize matmuls if not 1000s of person hours/years/etc of very "fine-grained" perf research? what the "abstraction" people don't understand (because they only deal in abstractions) is that achieving performance involves literally the antithesis of abstraction - you need to understand your hardware down to the gate level (sometimes).
> loop tiling, pipelining, shared memory swizzle, memory coalescing
have you ever applied any of these? the only way you could apply these as a generic (without consideration of your particular hardware) algo is using a tuner; this is of course widely the route taken but that's not an "understanding" of anything except guess and check.
"abstraction is the most important thing - look at pytorch it's the best framework because of the perfect/beautiful/brilliant abstractions" (re functorch or fx or dynamo).
ignoring entirely how much tedious and grueling bookkeeping/corner-casing/kernel-tuning (by a perpetual 100s of fulltime engineers) presenting such an "abstract" interface to the user requires.
Let's instead constructively talk about techniques in concrete items. If you look at OpenAI's Triton (which is also a small team of < 5 core contributors), what's this abstraction and their key to high performance? It's a tile-based programming model, where a tile could be conveniently lowered to vector instructions, coalesced memory access, and transformed to permuted layout. Its `dot` on tiles can be directly lowed to TensorCore-specific instructions. With those in design, without a huge team painfully maintaining the system, critical kernels like FlashAttention could be quickly developed within say 30 lines of code.
I know who you are and you should probably be out in the open with the fact that you have a conflict of interest in working at octo, a company that sells a very specific type of ML compiler.
>Let's instead constructively talk about techniques in concrete items. If you look at OpenAI's Triton
Pretty ironic you would call out Triton is being the right abstraction because while it is true philippe did a very good thing by moving things from warp level to block level, there is absolutely no one that thinks (myself included) that Triton is an abstraction.
Unfortunately, I don’t know much about you, and actually I don’t really think there is conflict of interest if you work in Modular, because Modular is also developing compiler abstractions, which is something I like and agree with, isn’t it? Let’s discuss about techniques, and it doesn’t have to be that heated :-)
To clarify, my point is matmuls can be solved with proper compiler abstractions, and it’s not that hard, and if you are working on a compiler, I believe you would more or less agree with that point, do you?
Liking Triton or not is a personal preference, and I use this as an example only because it’s gaining a lot of momentum at the moment, not saying it’s a perfect abstraction. If you personally don’t like it, I could also discuss about exo-compilation, tensor comprehension, but let’s always focus on concrete technical items :-)
https://halide-lang.org/tutorials/tutorial_lesson_05_schedul...
We recently have some follow-ups on this idea as well. Happy to discuss about it and I don't want to distract from this thread...
My comment is based on my personal experience: I did lead a 2nd/3rd grade undergrad to add software pipelining support and it worked within 1 month; we did get cutlass-level performance within 100 lines of code specifying the design space.
Writing assembly doesn’t scale across lots of platforms? Sure… the solution for matrix multiplication is to use the vendor’s BLAS.
If the vendor can’t at least plop some kernels into BLIS they don’t want you to use their platform for matmuls… don’t fight them.
The problem is already "solved" to almost everyone's satisfaction by being O(N), i.e. one optimized matrix math library per platform.
But if they can reduce that to O(1) by creating a tool that takes computing hardware characteristics (core/compute topology, instructions, memory heirarchy, ...), and outputs state-of-the-art optimized matrix multiply machine code, it would be a nice and useful result.
There's also the issue of poor documentation and learning material in the wild.
It's not fragmentation, they built a moat.
Sounds like they would oddly prefer memory latency to grow as least as fast as processing speeds, which would be terrible. Obviously, memory latency actually decreased, just not enough.
So it seems likely they made a mistake and actually meant that memory latency has decreased slower than processing speeds have increased, in other words, that it is not memory latency but memory random access throughput (which in rough approximation is about proportional to the inverse of memory latency) that has grown much slower than processing speeds.