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.
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.
#1 Best Overall
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
- Confirm that your GPU family and software stack appear in the official README.
- Install the package by following the installation section of that README for your CUDA or ROCm version.
- 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.
- Call
flash_attn_func(q, k, v). If your queries, keys, and values are stored as one packed tensor, callflash_attn_qkvpacked_funcinstead.
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.
Rank #2
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.
Quick wins for a faster PC:
Clear out junk files and repair common Windows errorsFree Scan →Scan for outdated or missing drivers - takes under a minuteDriver Scan →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.
Rank #3
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
- Install Triton and PyTorch following Triton’s current installation instructions, and confirm that the two versions work together on your GPU.
- Open the fused-attention tutorial and read the forward kernel first, then the backward kernel.
- 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.
- 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:
Rank #4
- 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.
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.
The Tool Desk
Outbyte PC Repair FREERepair Windows errors before they cause bigger problemsFix Now →Outbyte Driver Updater FREEScan for outdated or missing drivers - takes under a minuteDriver Scan →Quick Recap
| 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.

