FlexAttention: The Flexibility of PyTorch with the Performance of FlashAttention
pytorch.org
pytorch.org
We’re quite happy with this abstraction - happy to answer any questions about it!
Attention(Q,K,V) = Softmax(Q*K^T/sqrt(d_k))*V
FlexAttention seems to have found the right abstraction for the task.
That’s all without taking into account the broadcasting done on the batch dimension
Nice to see Pytorch already elegantly supporting this next step in research
/attention-gym/.venv/lib/python3.11/site-packages/torch/_subclasses/functional_tensor.py:258: UserWarning: Failed to initialize NumPy: No module named 'numpy' (Triggered internally at /Users/runner/work/pytorch/pytorch/pytorch/torch/csrc/utils/tensor_numpy.cpp:84.) cpu = _conversion_method_template(device=torch.device("cpu")) Traceback (most recent call last): File "/attention-gym/attn_gym/masks/document_mask.py", line 7, in <module> from torch.nn.attention.flex_attention import _mask_mod_signature ModuleNotFoundError: No module named 'torch.nn.attention.flex_attention'
It's very good. But note FlashAttention-3 is 1.5x - 2x faster than FlashAttention-2.
On Hopper, FlexAttention is currently about 80% of FlashAttention3's performance (about 500 TFLOPs peak)
Does anybody have a good starting point to learn with hands-on projects and also that could accommodate for flexattention?
A classifier for handwritten digits in the MNIST dataset is generally considered the "Hello World" of neural networks. I went over it in a course, but there are countless tutorials to be found online, i.e. https://www.digitalocean.com/community/tutorials/introductio...
Once you begin to understand how to handle data and how to define layers, you can start playing around with whatever your heart desires. The rabbit hole is vast and endless :)
Perhaps this tweet thread would be better.