Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

iTechGuides is reader-supported. When you buy through links on our site, we may earn an affiliate commission. As an Amazon Associate I earn from qualifying purchases. Learn more

To use FlashAttention-2 from PyTorch, you call a Python function from the official FlashAttention package on your query, key, and value tensors. Triton is not required for that. Triton appears in two separate places: the Triton project’s fused-attention tutorial is its own implementation of the same algorithm, written so you can read and modify it, and the package’s AMD ROCm support includes a Triton backend alongside a Composable Kernel one.

Three layers with one name

The phrase “FlashAttention-2 in Triton” mixes together three different things. Each one answers a different question, so it helps to read them separately.

Layer What it is What you do with it Primary source
Algorithm Exact attention designed around the GPU memory hierarchy and how work is partitioned across the GPU Understand the design and its trade-offs FlashAttention-2 paper
Triton tutorial One specific Triton implementation of the algorithm, including forward and backward paths and benchmark tables Read, run, and modify kernel code to learn how it is built Triton fused-attention tutorial
Official repository Python functions such as flash_attn_func and flash_attn_qkvpacked_func for scaled dot-product attention Import and call them from PyTorch code Official FlashAttention README

The tutorial and the package are separate artifacts. Do not assume they run identical code or support identical features.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

The paper frames its contribution as a change in work partitioning. Tri Dao, the paper’s author, writes: “We propose FlashAttention-2, with better work partitioning to address these issues.” The abstract names those issues as suboptimal partitioning across GPU thread blocks and warps.

Using FlashAttention-2 from PyTorch

This is the route most PyTorch users want. The official repository documents Python functions for scaled dot-product attention. Its documented options include causal attention, local windows, dropout, and ALiBi. Which of these work depends on the GPU, the backend, and the implementation path, so check the README before you depend on one.

Hardware and backend requirements

  • NVIDIA: the README lists Ampere, Ada, and Hopper GPU families. Its examples include A100, RTX 3090, RTX 4090, and H100.
  • AMD: the README describes ROCm support with Composable Kernel and Triton backends. Confirm which backend your ROCm version uses before relying on a specific feature.
  • Appearing on a supported-hardware list does not mean every feature behaves the same on every device.

Calling the function

  1. Confirm that your GPU family and software stack appear in the official README.
  2. Install the package by following the installation section of that README for your CUDA or ROCm version.
  3. Create the query, key, and value tensors on the GPU in a half-precision dtype that your installed version documents, such as fp16 or bf16.
  4. Call flash_attn_func(q, k, v). If your queries, keys, and values are stored as one packed tensor, call flash_attn_qkvpacked_func instead.
import torch
from flash_attn import flash_attn_func

# (batch, seqlen, nheads, headdim), fp16 or bf16, on the GPU
q = torch.randn(2, 1024, 8, 64, device="cuda", dtype=torch.float16)
k = torch.randn_like(q)
v = torch.randn_like(q)

out = flash_attn_func(q, k, v, causal=True)
print(out.shape)  # torch.Size([2, 1024, 8, 64])

The argument names and tensor layout above follow the package’s usual signature. Confirm both against the README for the exact version you install before copying the snippet into a model.

The Triton tutorial: what it is and how to use it

Triton is a Python-based language for writing GPU kernels. The fused-attention tutorial presents its sample as an implementation of FlashAttention-2. The tutorial states: “This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao.” The page includes both forward and backward paths, along with benchmark tables.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

The benchmark tables reflect the hardware and settings used on that page at one point in time. The page can change as the main branch changes, so read the tables as one configuration rather than a fixed result.

When the tutorial is the right tool

  • You want to see how the algorithm’s work partitioning maps onto block-level Triton code.
  • You want to change the kernel itself, for example to experiment with a different masking or tiling scheme.
  • You want to study the backward pass as well as the forward pass, since the tutorial includes both.

Running and modifying it

  1. Install Triton and PyTorch following Triton’s current installation instructions, and confirm that the two versions work together on your GPU.
  2. Open the fused-attention tutorial and read the forward kernel first, then the backward kernel.
  3. Run the tutorial on your own GPU. Compare its output with the page’s benchmark tables only as a rough reference, because the hardware and settings will differ.
  4. Change one thing at a time, such as the causal flag or a block size, and check the result against a standard attention computation on small inputs before you measure speed.

Treat the tutorial as a reference to learn from. Its feature set, tuning, and hardware coverage are not stated to match the package, so do not assume they do.

Choosing a path

Goal Start with Check before relying on it
Use attention in a PyTorch model on NVIDIA GPUs Official package functions Your GPU family in the README, and whether the options you need (causal, window, dropout, ALiBi) are supported on your path
Use attention on AMD GPUs The package’s ROCm support Which backend (Composable Kernel or Triton) your ROCm version uses
Learn how FlashAttention maps to GPU code Triton fused-attention tutorial That the page matches the version of the code you are reading
Experiment with a modified attention kernel The tutorial, copied and changed Correctness against a reference first, then speed on your own GPU

Why FlashAttention-2 is faster

FlashAttention reduces memory traffic, meaning the amount of data moved between levels of GPU memory. FlashAttention-2 keeps that approach and improves how the work is partitioned, which the paper identifies as the limit on earlier performance. The paper names three core changes:

  • Fewer non-matmul FLOPs. It reduces floating-point operations that are not matrix multiplications.
  • More parallelism. It parallelizes attention across thread blocks even for a single attention head, so one head’s work can be spread across the GPU.
  • Less inter-warp communication. It reduces communication between warps through shared memory.
Independent reader supportYour contribution helps us test, update, and keep practical guides available for everyone.Support on Ko-Fi

What the paper measured

The paper, by Tri Dao, reports benchmarks on an A100 80GB SXM4 with sequence lengths from 512 to 16k, a hidden dimension of 2048, and head dimensions of 64 or 128. The paper dates from 2023, so these are historical results for that setup.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.
Reported figure Comparison What it measures
1.3–2.5× faster Against FlashAttention in Triton, across the paper’s evaluated comparisons Attention computation on the paper’s setup
Forward about 1.3–1.5×; backward about 2× The paper’s individual forward and backward comparisons Attention computation, per direction
Up to 10× faster Against a standard attention implementation in PyTorch, in the paper’s evaluated comparisons Attention computation on the paper’s setup
Up to 230 TFLOPs/s, 73% of theoretical maximum on A100 FlashAttention-2 kernel throughput Attention kernel throughput
Up to 225 TFLOPs/s and 72% model FLOPs utilization per A100 Reported end-to-end training experiments Full training throughput, not a single kernel
  • Kernel-level and end-to-end figures measure different things. Do not add them together or treat one as a substitute for the other.
  • These are experimental results for the paper’s setup. They are not guarantees and not a current leaderboard.
  • The paper does not measure the RTX 4090 or other GPUs. Do not read the A100 figures as predictions for another card.

Common mistakes

  • Assuming Triton is required. It isn’t. The package’s Python functions are the direct path from PyTorch.
  • Reading tutorial benchmark tables as package performance. They describe the tutorial’s configuration on the page.
  • Assuming feature parity. NVIDIA, AMD, and each backend path can differ in which options they support.

Troubleshooting

Symptom Likely cause What to do
Installation or import fails GPU family or CUDA/ROCm version is outside the README’s support, or software versions do not match Compare your GPU family and software stack with the README, then reinstall following the section for your version
Output differs from standard attention Different dtype, a causal setting that doesn’t match, or a tolerance that is too strict for half precision Re-run on small inputs in the same dtype, confirm the causal setting matches, and use a tolerance suited to fp16 or bf16
An option such as dropout or a local window does not behave as expected The option is not supported on your backend or implementation path Check the README’s support for your backend and path
Tutorial code differs from the documentation you read The main branch has changed since you read the page Pin a specific commit or tag and read the documentation for that version
Slower than the paper’s numbers Different GPU, dtype, sequence length, or head dimension, or a kernel figure compared with an end-to-end one Compare like with like: same GPU class, dtype, and workload shape, and the same kind of measurement

Product prices and availability are accurate as of the date/time indicated and are subject to change. Any price and availability information displayed on Amazon at the time of purchase will apply.