PyTorch Native Architecture Optimization: Torchao
pytorch.org
pytorch.org
TF's doesn't seem very good. I just tried to figure out how to learn a linear mapping with TF and went through this:
1. googled "linear layer in tensorflow" and got to the page about linear.
2. spent 5 minutes trying to understand why monotonicity would be a central tenet of the documentation
3. realizing that's not the right "linear" I couldn't think of what the appropriate name would be
4. I know MLPs have them, google "tensorflow mlp example"
5. click the apr '24 page: https://www.tensorflow.org/guide/core/mlp_core
6. read through 10[!] code blocks that are basically just boiler-plate setup of data and visualizations. entirely unrelated to MLPs
7. realize they call it "dense" in tensorflow world
8. see that "dense" needs to be implemented manually
9. think that's strange, google "tensorflow dense layer"
10. find a keras API (https://www.tensorflow.org/api_docs/python/tf/keras/layers/D...)
I have seen some good ones, too, of course.
(This pattern is relatively easy to understand: smart people creating something get their gratification from the creation process, not writing tedious documentation; and this is systemically embedded for people at Google, who are probably directly incentivised in a similar way.)
JAX is right there. No need to beat a dead horse when there's a stallion in the stables.
Agreed of course but it's not like they came up with this approach from scratch. They seem to have just picked it up from Theano (now Aesara/PyTensor).
From what I can tell Google is moving in a direction that doesn't require tensorflow, and I don't see it gaining signficant adoption outside google, so it seems most likely we will simply see it deprecated in about 10 years. It's best to see it as a transitional technology that Jeff Dean created to spur ML development internally, which was mistakenly open sourced, and now, Jeff's reports typically use Jax or other systems.
Essentially the latency overhead comes from quantizing and dequantizing weights and activations. For large layers this overhead is small because by quantizing your weights for example you reduce memory bandwidth pressure but for small layers the overhead of potentially looking up a table, reading scaling factors, quantization/dequantization and finally handling zero points might not be worth it.
However, even if such overhead exists you can still quantize your model and get it to be smaller it might not be faster is the problem. We solve the speed problem in 2 ways - `torch.compile()` will fuse operations like a dequant and matmul into a single kernel and `torchao.autoquant()` will do kernel level profiling to see whether a layer is actually made faster when quantizing and if not it skips quantizing that layer.
In these cases the only path forward we have is writing custom Metal kernels and plugging those in. That work is still ongoing and we'll hopefully have more to share soon.
Granted after more upfront effort compilers are just such a significant UX boost that indeed you are making me question why I don't spend more time working on this myself lol
Basically PyTorch is a large library where CI takes a long time to run which means merging code is hard and adding new dependencies is challenging and there are stringent constraints on BC breaking changes
Instead what torchao did and many other repos like torchtune, torchchat, torchtitan did was move out of core and it helps keep the core PyTorch library leaner with a smaller binary size and it really lets the team "out of core" focus on optimizing for their needs
Unfortunately the argument for what gets better changes over time, for example torch.compile initially a new repo called torchdynamo was built out of core to move fast but eventually merged back because everyone wanted to use it. Now torch.compile dev velocity is still quite fast and so now we have to tell people to use nightlies instead of official stable releases to which some people have asked me why don't you move torch.compile out of core
My 2c is the ecosystem will be much stronger and teams can move faster if they develop out of core so that's the tradeoff we picked for torchao. We managed to for example merge a few custom CPP kernels like fp6 or Marlin that would have challenging to motivate in core since those are still quite experimental and need to stand the test of time.
But we have had quantization algorithm developers such as HQQ or Autoround merge their code in to get composability and serialization for free. We view quantization algorithms as the top layer and going down you have quantized tensors, quant primitives like dequant/quant and finally basic dtypes like uint1-7 and float3-8. Personally why I spent so much time on AO was I was hoping we could make it easier for people to express their quantization algorithms in easy to read PyTorch code and if they must use custom kernels we also have some tutorials for how to integrate custom cuda and triton ops.
Most of those discussions have been happening on #torchao on discord.gg/gpumode so if you need to chat back and forth feel free to reach out to the team there otherwise Github also works.
A minor nitpick on the copy (and even then, it might just be me): I find "97% speedup" and "50% speedup" really hard to parse — a "30x speedup" or "97% reduction of time taken" immediately tell me what is being achieved!
Great results once I get my head around them, though!
I guess you are right and it's probably the latter, but obviously better language would have avoided any doubt.
But that's waiting for Blackwell to be released so we get the hardware support. SO recommendation for now would be to use either fp8 training or int8 training
AFAIU int4 matrix multiplication is supported by cuda, but I'm not sure about other operations. The blog post mentioned fp6, and I don't think this is supported by cuda. Or maybe the data are upscaled to something common like fp16 before doing math?
Will this let me use uint8 arrays as indexing arrays? A problem I have is that pytorch forces me to use uint64 for fancy indexing.
https://github.com/google/aqt is more explicit and preferable IMO.
Neither are as user-friendly as what Torchao has presented here.