Table of Contents
- What I am building
- Step 1: the tiled forward in PyTorch
-
A primer to
tl.make_block_ptrbaseas the starting point in memory of a tensorshapetells Triton how big the tensor isstridestells Triton how far apart neighbours are in memoryoffsetstells Triton where the tile startsblock_shapetells Triton the size of a tileordertells Triton which axis is side by side in memory.advance()tells Triton how to move the tileboundary_checktells Triton what happens past the edge- Put it all together
- Cheat sheet
- Step 2: the Triton forward kernel
- Step 3: testing and bugs found
- Step 4: the benchmark, and the baseline number
- What is wrong with this kernel
- Future articles
I wanted to learn kernel optimisation by starting with Triton. While I was able to follow basic tutorials like vector addition, the learning curve increased sharply when I moved towards more complex algorithms like Flash Attention 2.
After some exploration of various tutorials, I settled on following Stanford CS336 assignment 2 which provides the implementation of Flash Attention 2 (FA2) as an assignment. I liked it for the following reasons:
- It began by asking for a PyTorch implementation of FA2, which helped to solidify my understanding of FA2 before introducing complexities associated with Triton.
- It made assumptions to simplify the implementation of FA2 to help me build up my understanding of basic concepts in Triton first. These assumptions included (i) inputs of shape
(batch, seq, d)with no head dimension; (ii) tile sizes of at least 16x16, which I hard-coded to 16x16; and (iii) sequence lengths and head dimensions that are powers of 2 and at least 16, so out-of-bounds accesses could be ignored.
Future articles will relax these assumptions.
I also retained some of the mistakes I made when implementing FA2, along with explanations of these mistakes and their corrections, because I learnt the most from mistakes.
All numbers in this series are from a single RTX 4070 Super (Ada, sm_89, 12 GB, roughly 100 KB of shared memory per SM) under Linux.
This article assumes familiarity with PyTorch and Einops, and no prior experience with Triton. This article was written with the assistance of AI.
If you spot any mistakes in this article, please let me know. I have also written articles on online softmax and FlashAttention by hand to build intuition for these two algorithms.
What I am building
The FlashAttention-2 (FA2) forward pass: for each query tile Q_i, iterate over key and value tiles K_j and V_j, rescaling the running output O_i on every step and dividing by l only once at the end (FA1 re-normalised at every step). The outer loop goes through the query tiles; on the GPU, each query tile gets its own program, and all the programs run in parallel.
The kernel computes standard scaled-dot-product attention without materialising the N×N score matrix. For a query tile Q_i of B_q rows and a key tile K_j of B_k rows, it computes the B_q × B_k score block, folds it into a running softmax, and moves to the next key tile. When all key tiles have been visited, the output rows for Q_i are complete and are written once.
Three running quantities per query tile make that possible, initialised on line 6:
-
m_i, the running maximum score for each row (line 10) -
l_i, the running row sum of exp(S_i - m_i) (line 12) -
O_i, the running sum of exp(S_i - m_i) times the value rows, kept unnormalised (line 13)
The superscript (j) marks the value after the j-th key tile.
On every step the previous l_i^(j-1) and O_i^(j-1) are rescaled by exp(m_i^(j-1) - m_i^(j)), which is 1 when the maximum is unchanged, before the new tile's contribution is added. This is the online softmax, which is applied one tile at a time.
At the end of each query tile, the kernel also writes a second output L, which is computed on line 16 and stored on line 18. L is the per-row logsumexp of the scaled scores L_i = m_i + log(l_i), which equals log Σ_k exp(S_ik) over the whole row. The backward pass will use it to recompute the softmax probabilities tile by tile as P = exp(S - L), since exp(S - m) / l = exp(S - m - log(l)) = exp(S - L). Storing L instead of P is what lets FA's backward avoid the NxN matrix too. The backward will be covered in a later article. For now L is computed, stored, and tested, but not used.
Step 1: the tiled forward in PyTorch
Before diving into Triton, I coded FA2 with PyTorch to familiarise myself with the algorithm. The PyTorch implementation is slow by design because it has a Python double loop over tiles, and launches a handful of small CUDA ops per iteration. The goal here is to be correct, not fast.
# flashattention_autograd_function_pytorch.py
import math
import torch
import einops
class FlashAttentionPytorch(torch.autograd.Function):
@staticmethod
def forward(ctx, Q, K, V, is_causal=False):
# Note: Tile size is fixed at 16 as a simplifying assumption
tile_size = 16
# Split the sequence dimension into tiles: (..., T, B, d)
# Leading dims are arbitrary (batch, heads, ...). The sequence axis N is split
# into Tq tiles of Bq rows, so N must be a multiple of tile_size.
Q_t = einops.rearrange(Q, "... (Tq Bq) d -> ... Tq Bq d", Bq=tile_size)
K_t = einops.rearrange(K, "... (Tk Bk) d -> ... Tk Bk d", Bk=tile_size)
V_t = einops.rearrange(V, "... (Tk Bv) d -> ... Tk Bv d", Bv=tile_size)
O = torch.empty_like(Q)
L = torch.empty(Q.shape[:-1], device=Q.device, dtype=Q.dtype)
scale = 1.0 / math.sqrt(Q.shape[-1])
for i in range(Q_t.shape[-3]): # outer loop: query tiles
Q_i = Q_t[..., i, :, :]
O_i = torch.zeros_like(Q_i)
l_i = torch.zeros(Q_i.shape[:-1] + (1,), device=Q.device, dtype=Q.dtype)
m_i = torch.full(Q_i.shape[:-1] + (1,), -torch.inf, device=Q.device, dtype=Q.dtype)
for j in range(K_t.shape[-3]): # inner loop: key tiles
K_j = K_t[..., j, :, :]
V_j = V_t[..., j, :, :]
S_ij = einops.einsum(Q_i, K_j, "... Bq d, ... Bk d -> ... Bq Bk") * scale
m_new = torch.maximum(m_i, S_ij.amax(dim=-1, keepdim=True))
P_ij = torch.exp(S_ij - m_new)
alpha = torch.exp(m_i - m_new) # rescale factor for the old state
l_i = alpha * l_i + P_ij.sum(dim=-1, keepdim=True)
O_i = alpha * O_i + einops.einsum(P_ij, V_j, "... Bq Bk, ... Bk d -> ... Bq d")
m_i = m_new
O[..., i * tile_size:(i + 1) * tile_size, :] = O_i / l_i
L[..., i * tile_size:(i + 1) * tile_size] = (m_i + torch.log(l_i)).squeeze(-1)
ctx.save_for_backward(Q, K, V, O, L)
ctx.is_causal = is_causal
return O, L
A few points to note:
-
m_istarts at negative infinity so that the first tile's row maximum "wins" in line 10 of the algorithm above.m_istarting at negative infinity also ensures on the first stepalpha = exp(-inf - m_new)is exactly 0, not NaN, so that the rescale in lines 12 and 13 runs unchanged and there is no special case to handle for the first iteration. -
O_iis the weighted sum of value rows, and the softmax needs that sum divided by the total weightl_i. There are two ways to book-keep: (i) keep the sum and divide once at the end (FA2); or (ii) keep the average and re-divide every time a new tile arrives (FA1). FA2's approach saves two elementwise passes over aB_q x dtile on every inner loop. -
Lis stored, notmandlseparately. The backward pass recomputes each tile ofSfromQandK, then recovers that tile's probabilities asexp(S - L). For context, assumingN = 4096,Swith dimensionNxNper head per batch element will have about 16.8 million entries, or about 33MB in bf16.Lwith dimensionNx1per head per batch element contains 4096 entries, or about 16KB. - The
is_causalflag is accepted and ignored. The PyTorch version does not implement causal masking.
A primer to tl.make_block_ptr
Before I introduce the implementation of FA2 in Triton, following Stanford CS336 assignment 2, I will introduce tl.make_block_ptr.
tl.make_block_ptr creates a block pointer. A block pointer is a tile sliding over a tensor. The code to create a block pointer in Triton is below. Using x_block_ptr, the kernel below adds up each row of x, with each column multiplied by a weight. The kernel is called weighted_sum_fwd. Instead of loading the entire row at once, the kernel slides across the columns of x one tile at a time, in a loop.
x_block_ptr = tl.make_block_ptr(
x_ptr,
shape=(NUM_ROWS, D),
strides=(x_stride_row, x_stride_dim),
offsets=(row_tile_idx * ROWS_TILE_SIZE, 0),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
)
# Walk across the columns of x, one tile per step.
# cdiv divides and rounds up, so no column is missed.
for i in range(tl.cdiv(D, D_TILE_SIZE)):
row = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
... # multiply by weight, add to the totals
x_block_ptr = x_block_ptr.advance((0, D_TILE_SIZE)) # move one tile right
There are 6 arguments in tl.make_block_ptr, illustrated in the image below. I will go through them one by one. The illustrations and HTML visualisations are drawn by Claude Opus 5.5.
The blue rectangle is the tile, i.e. the part of the tensor the block pointer points at. It sits over part of the tensor: .advance() slides it, and tl.load reads what is inside. Triton calls it a block, which is where make_block_ptr and block_shape get their names.
Why use a block pointer
A Triton kernel runs many programs in parallel. Each program works on a small tile of data. To read a tile, a program needs the memory address of every element in the tile. It also must not read past the edges of the tensor. There are two ways to do this.
The classical way with pointer maths
Before block pointers, programmers built the addresses directly, and many codebases, including the official Triton tutorials, still do. Below is the weighted_sum_fwd kernel written the classical way. I will explain how to read Triton code written the classical way as it is still a useful skill to master.
import triton
import triton.language as tl
@triton.jit
def weighted_sum_fwd(
x_ptr, weight_ptr, # inputs: where x and weight start in memory
output_ptr, # output: where the results go
x_stride_row, x_stride_dim, # x is 2D: one stride per axis
weight_stride_dim, # weight is 1D: one stride (usually 1)
output_stride_row, # output is 1D: one stride (usually 1)
NUM_ROWS, D, # the size of x
ROWS_TILE_SIZE: tl.constexpr, # tile sizes, fixed when the kernel is compiled
D_TILE_SIZE: tl.constexpr,
):
row_tile_idx = tl.program_id(0)
# 1. Which rows this program owns
rows = row_tile_idx * ROWS_TILE_SIZE + tl.arange(0, ROWS_TILE_SIZE)
row_mask = rows < NUM_ROWS
output = tl.zeros((ROWS_TILE_SIZE,), dtype=tl.float32)
for i in range(tl.cdiv(D, D_TILE_SIZE)):
# 2. Which columns this step covers
cols = i * D_TILE_SIZE + tl.arange(0, D_TILE_SIZE)
col_mask = cols < D
# 3. A 2D grid of addresses, built by broadcasting
x_ptrs = (x_ptr + rows[:, None] * x_stride_row
+ cols[None, :] * x_stride_dim)
w_ptrs = weight_ptr + cols * weight_stride_dim
# 4. Masks for the edges, combined by hand
row = tl.load(x_ptrs, mask=row_mask[:, None] & col_mask[None, :], other=0.0)
weight = tl.load(w_ptrs, mask=col_mask, other=0.0)
output += tl.sum(row * weight[None, :], axis=1)
tl.store(output_ptr + rows * output_stride_row, output, mask=row_mask)
The code to start the kernel in Python is below, i.e. the launch code. grid says how many programs to run, and the other values fill the signature of weighted_sum_fwd in order. The block pointer version later in this section has the same signature, so the same launch code works for both.
import torch
import triton
x = torch.randn(5, 10, device="cuda") # NUM_ROWS = 5, D = 10
weight = torch.randn(10, device="cuda")
output = torch.empty(5, device="cuda")
NUM_ROWS, D = x.shape
ROWS_TILE_SIZE, D_TILE_SIZE = 2, 4
# One program per ROWS_TILE_SIZE rows: 5 rows need cdiv(5, 2) = 3 programs
grid = (triton.cdiv(NUM_ROWS, ROWS_TILE_SIZE),)
weighted_sum_fwd[grid](
x, weight, output, # a tensor is passed as its start address
x.stride(0), x.stride(1), # 10, 1
weight.stride(0), # 1
output.stride(0), # 1
NUM_ROWS, D,
ROWS_TILE_SIZE=ROWS_TILE_SIZE, D_TILE_SIZE=D_TILE_SIZE,
)
What each line does
I follow the case from the launch code above. x has NUM_ROWS = 5 rows and D = 10 columns, stored row by row. So x_stride_row = 10 (one row down skips 10 elements) and x_stride_dim = 1 (one column right is the next element). The kernel also takes one stride for each 1D tensor, weight_stride_dim and output_stride_row. Both are 1, because the elements of each vector sit next to each other in memory. Tiles are ROWS_TILE_SIZE = 2 rows by D_TILE_SIZE = 4 columns. For the rest of this section, I will trace program 2, which handles the last row.
In one sentence: each program adds up x[r, c] * weight[c] over all columns c, for each of its rows r. It works 4 columns at a time and writes one total per row at the end.
Setup
row_tile_idx = tl.program_id(0)
Because there are 5 rows, and each program handles 2 rows, 5 rows need 3 programs (5 ÷ 2, rounded up), hence grid = 3. The programs are numbered 0, 1 and 2. When the kernel runs, the GPU gives each program its own number, and tl.program_id(0) reads it. The 0 asks for the first axis of this numbering. Here grid has only one entry, so there is only one axis. The program I am tracing gets 2, so it handles rows 4 and 5. It is the last program, and its tile is only half full: row 4 exists, but row 5 does not. Every line below runs separately inside each program, and the programs may run at the same time, in any order.
1. Which rows this program owns
rows = row_tile_idx * ROWS_TILE_SIZE + tl.arange(0, ROWS_TILE_SIZE)
tl.arange(0, 2) makes the small vector [0, 1]. It is a Triton tensor that holds both numbers at once. The line computes 2 * 2 = 4 and adds it to each element, so rows = [4, 5]. These are the rows this program handles. Triton needs to know the length of this vector when it compiles the kernel, which is why ROWS_TILE_SIZE is marked tl.constexpr, which means it is a value fixed at compile time.
row_mask = rows < NUM_ROWS
This compares each element with 5: [4 < 5, 5 < 5] gives [True, False]. This is a mask: one True or False per element, where True means "this spot exists, so it is safe to use". Row 5 does not exist, because x only has rows 0 to 4. This line runs once before the loop because the rows do not change.
output = tl.zeros((ROWS_TILE_SIZE,), dtype=tl.float32)
This makes one running total per row, starting at [0.0, 0.0]. It uses 32-bit floats even if x holds 16-bit numbers, because adding many low-precision numbers builds up rounding error. You can try the CS336 Assignment 2 experiments with floating points to understand why accumulators need to have higher precision.
for i in range(tl.cdiv(D, D_TILE_SIZE)):
tl.cdiv divides and rounds up: cdiv(10, 4) = 3, so i takes the values 0, 1 and 2. Rounding up matters because plain division would give 2 steps and skip columns 8 and 9. Unlike the rows and row_mask vectors above, this is a normal loop where the steps run one after another. The for loop is not parallelised because each step adds to the totals left by the step before.
2. Which columns this step covers
cols = i * D_TILE_SIZE + tl.arange(0, D_TILE_SIZE)
The same pattern as rows, but across the columns: [0, 1, 2, 3] when i = 0, then [4, 5, 6, 7] when i = 1, then [8, 9, 10, 11] when i = 2. I will walk through the case when i = 2, because that is where the edge of the tensor matters.
col_mask = cols < D
This gives [True, True, False, False]. Columns 10 and 11 do not exist, hence are masked as False.
3. A 2D grid of addresses
x_ptrs = (x_ptr + rows[:, None] * x_stride_row
+ cols[None, :] * x_stride_dim)
tl.load needs one address for each of the 8 elements in the tile. This line builds all 8 at once, from rows = [4, 5] (step 1) and cols = [8, 9, 10, 11] (step 2).
The strides turn row and column numbers into distances in memory. Each row of x holds 10 elements, so row 4 starts 4 × 10 = 40 elements from the start of x, and row 5 starts at 50.
Each column adds 1 more element (x_stride_dim = 1). rows[:, None] makes the rows a column, and cols[None, :] makes the columns a row. On their own, rows and cols are just flat lists of numbers:
rows = [4, 5]
cols = [8, 9, 10, 11]
I want one address for every combination of a row and a column: 2 rows × 4 columns = 8 addresses. Think of an addition table, with rows down the side, columns across the top, and each cell adding its row's number to its column's number.
-
rows[:, None]stands the list up as a column (2 rows, 1 column):
[[4],
[5]]
-
cols[None, :]lays the list down as a row (1 row, 4 columns):
[[8, 9, 10, 11]]
The : means "keep all the values," and None means "add a new direction here, with size 1." Where you put the None decides whether the list stands up or lies down.
When you add a column to a row, Triton copies the column across and the row down until both are 2 × 4. Then it adds cell by cell. That's the addition table:
[[40], + [[8, 9, 10, 11]] → [[48, 49, 50, 51],
[50]] [58, 59, 60, 61]]
Broadcasting stretches both to a 2 × 4 shape and adds them cell by cell (panel A in the picture below):
| col 8 | col 9 | col 10 | col 11 | |
|---|---|---|---|---|
| row 4 (starts at 40) | 48 | 49 | 50 | 51 |
| row 5 (starts at 50) | 58 | 59 | 60 | 61 |
Each number counts elements from x_ptr, the address of x[0, 0]. So x_ptr + 48 is the address of x[4, 8]: 40 to reach row 4, plus 8 to reach column 8. Triton counts in elements, not bytes.
w_ptrs = weight_ptr + cols * weight_stride_dim
The weight tensor is a single list with one number per column of x, and every row of x uses the same weights. So only the column positions matter: four addresses, weight_ptr + [8, 9, 10, 11]. No rows are involved, so no grid is needed.
4. Loading safely at the edges
row = tl.load(x_ptrs, mask=row_mask[:, None] & col_mask[None, :], other=0.0)
First, the two masks are combined into a grid the same way as the addresses. row_mask becomes a column, col_mask becomes a row, and & ("and") makes a cell True only if both of its parts are True. This is panel B:
[[True, True, False, False],
[False, False, False, False]]
tl.load then reads memory only at the addresses where the mask is True. Where it is False, it does not touch memory and uses the value given by other (here 0.0) instead. The result is the 2 × 4 tile in panel C: [[x[4,8], x[4,9], 0, 0], [0, 0, 0, 0]].
The masking matters because x has 5 × 10 = 50 elements, at addresses 0 to 49. Every False cell is at address 50 or more, past the end of x. Reading there returns garbage values or crashes the kernel.
Program 2 on its last loop step (i = 2). A: the 8 addresses, built from rows and cols. B: the mask, built the same way from row_mask and col_mask. C: the tile that tl.load returns.
weight = tl.load(w_ptrs, mask=col_mask, other=0.0)
This gives [weight[8], weight[9], 0, 0].
output += tl.sum(row * weight[None, :], axis=1)
weight[None, :] makes the weights a row, so they multiply every row of the tile. The product has dimension 2 × 4. tl.sum(..., axis=1) adds across the columns, which are axis 1, giving one number per row. Those numbers are added to the running totals in output.
The 0s at the edges add nothing, so the half-empty tile gives the right answer. Row 5 is all 0s, so its total stays 0.
After the loop
tl.store(output_ptr + rows * output_stride_row, output, mask=row_mask)
After the last step, output holds one total per row: [total for row 4, total for row 5]. There are no columns left, because tl.sum(..., axis=1) added them up on every step.
The output tensor is a single list with one number per row of x, so each total needs just one address. With output_stride_row = 1, rows * output_stride_row is [4, 5], so the addresses are output_ptr + [4, 5].
tl.store then matches up three lists of 2, position by position:
| first | second | |
|---|---|---|
| address | output_ptr + 4 |
output_ptr + 5 |
| value | total for row 4 | total for row 5 |
row_mask |
True |
False |
It writes row 4's total, which is the sum of x[4, c] * weight[c] over all 10 columns. It skips row 5, which would be past the end of output.
Unlike in step 3, rows is not turned into a column here. [:, None] and [None, :] are only needed to combine two lists into every combination, like rows × columns. Here there is only one list, so matching values one to one is exactly right.
Each program writes only its own rows, so no two programs ever write to the same place.
What you have to get right with pointer maths
With pointer maths, every detail is your job:
- the index vectors, made with
tl.arange, - the broadcasting, with
[:, None]and[None, :]on the right vectors, - the strides, each multiplied into the right index,
- one mask per axis, combined with
&, plus anothervalue for the spots the mask leaves out, - new column indices on every loop step.
A slip in any of these gives wrong numbers, and there may not be any error message.
The block pointer way
Here is the same kernel written with block pointers. The signature is the same, so the launch code above works unchanged. Three block pointers are made once, before the loop. After that, three calls do the rest: tl.load reads a tile, .advance() moves it, and tl.store writes a tile.
import triton
import triton.language as tl
@triton.jit
def weighted_sum_fwd(
x_ptr, weight_ptr, # inputs: where x and weight start in memory
output_ptr, # output: where the results go
x_stride_row, x_stride_dim, # x is 2D: one stride per axis
weight_stride_dim, # weight is 1D: one stride (usually 1)
output_stride_row, # output is 1D: one stride (usually 1)
NUM_ROWS, D, # the size of x
ROWS_TILE_SIZE: tl.constexpr, # tile sizes, fixed when the kernel is compiled
D_TILE_SIZE: tl.constexpr,
):
row_tile_idx = tl.program_id(0)
# 1. Describe each tensor, and where this program's tile starts
x_block_ptr = tl.make_block_ptr(
x_ptr,
shape=(NUM_ROWS, D),
strides=(x_stride_row, x_stride_dim),
offsets=(row_tile_idx * ROWS_TILE_SIZE, 0),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
)
weight_block_ptr = tl.make_block_ptr(
weight_ptr,
shape=(D,),
strides=(weight_stride_dim,),
offsets=(0,),
block_shape=(D_TILE_SIZE,),
order=(0,),
)
output_block_ptr = tl.make_block_ptr(
output_ptr,
shape=(NUM_ROWS,),
strides=(output_stride_row,),
offsets=(row_tile_idx * ROWS_TILE_SIZE,),
block_shape=(ROWS_TILE_SIZE,),
order=(0,),
)
output = tl.zeros((ROWS_TILE_SIZE,), dtype=tl.float32)
for i in range(tl.cdiv(D, D_TILE_SIZE)):
# 2 to 4. Columns, addresses and edge masks: all handled inside tl.load
row = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
weight = tl.load(weight_block_ptr, boundary_check=(0,), padding_option="zero")
output += tl.sum(row * weight[None, :], axis=1)
# Move both tiles one tile to the right for the next step
x_block_ptr = x_block_ptr.advance((0, D_TILE_SIZE))
weight_block_ptr = weight_block_ptr.advance((D_TILE_SIZE,))
tl.store(output_block_ptr, output, boundary_check=(0,))
Triton does the index maths, the broadcasting and the masks for you. The code is not shorter here, because the make_block_ptr calls are long. However, block pointers help to abstract away the handling of memory addresses, which is error-prone.
Each piece of the hand-written version has a matching argument:
| By hand | With a block pointer |
|---|---|
x_ptr |
base |
rows < NUM_ROWS, cols < D
|
shape plus boundary_check
|
* x_stride_row, * x_stride_dim
|
strides |
row_tile_idx * ROWS_TILE_SIZE + ..., i * D_TILE_SIZE + ...
|
offsets, then .advance()
|
tl.arange(0, ROWS_TILE_SIZE), tl.arange(0, D_TILE_SIZE)
|
block_shape |
[:, None] and [None, :]
|
done for you |
other=0.0 |
padding_option="zero" |
| nothing |
order, an extra hint about memory layout |
Two lines are exactly the same in both versions: row_tile_idx = tl.program_id(0) at the start, and the maths, output += tl.sum(row * weight[None, :], axis=1). For the same inputs, both versions give the same result.
The kernel makes three block pointers: one for x (2D), one for weight (1D), and one for output (1D). The rest of this guide takes the arguments one at a time, using x since it is 2D and hence more general than 1D.
base as the starting point in memory of a tensor
GPU memory is a long line of slots and not a grid. A 2D tensor is stored on that line row by row, i.e. all of row 0, then all of row 1, and so on.
The number in each cell is its slot in memory. Row 0 fills slots 0 to 3, row 1 fills slots 4 to 7, and so on.
The first argument, base, is the address of the very first element, x[0, 0]. On its own, it says nothing about the size or layout of the tensor. The other arguments add that.
In the kernel: x_ptr, weight_ptr and output_ptr.
shape tells Triton how big the tensor is
shape is the size of the full tensor, not the size of the tile. For x it is (NUM_ROWS, D). For weight it is (D,), and for output it is (NUM_ROWS,).
The thick outline is shape. The striped red cells are outside the tensor. The last tile in a row can hang over the edge, like this one.
Triton uses shape to know where the tensor ends. NUM_ROWS and D do not have to be multiples of the tile size, so a tile can stick out past the edge. When you load with boundary_check, Triton compares each spot in the tile with shape and fills the outside spots with a padding value, such as 0. The edges section below shows this.
shape can be a normal runtime number; it does not need to be a tl.constexpr (a value fixed at compile time).
strides tells Triton how far apart neighbours are in memory
strides has one number per axis. Each number says how many slots you jump in memory when you take one step along that axis.
Try the three layouts. The grid is the same every time, but the elements land in different memory slots. The strides are what tell Triton which layout it is looking at.
Strides count elements, not bytes. PyTorch's x.stride() gives you exactly these numbers, so the Python side passes x.stride(0) and x.stride(1).
offsets tells Triton where the tile starts
offsets is the (row, column) of the tile's top-left element. It counts elements.
Each program (one copy of the kernel running on the GPU) handles its own band of rows. tl.program_id(0) tells it which band. To turn a band number into a row number, multiply by the band height: row_tile_idx * ROWS_TILE_SIZE.
The column offset is 0 because every program starts at the left edge. The loop then walks it to the right.
The other two pointers follow the same idea. weight starts at (0,). output starts at (row_tile_idx * ROWS_TILE_SIZE,), so each program writes exactly the rows it read. Offsets can be computed at run time, and usually are.
block_shape tells Triton the size of a tile
block_shape is the size of the tile. tl.load always gives you a tensor of exactly this shape, even when part of the tile is past the edge.
Triton compiles the kernel for one fixed tile size, so these numbers must be known at compile time. That is why the kernel marks ROWS_TILE_SIZE and D_TILE_SIZE as tl.constexpr. Triton requires each size to be a power of 2. Triton splits a tile evenly across the GPU's threads, which come in groups sized by powers of 2, so this keeps the split even. Tensors can be any size, and tiles at the edges are padded.
Bigger tiles mean fewer loop steps, but each program needs more registers (the small, fast storage inside the GPU chip, which runs out quickly). Picking good sizes is an empirical question, answered by profiling.
order tells Triton which axis is side by side in memory
Walk through memory one slot at a time and watch each element's position. In a matrix stored row by row, the column number changes at every step, while the row number changes only once per row. So the column axis changes "fastest" and the row axis "slowest". order lists the axes from fastest to slowest, so here it is (1, 0).
Each stop shows the memory slot of that element. Solid lines connect neighbours in memory. Dashed lines jump to the next row (left) or the next column (right).
For a normal row-major matrix, elements in a row sit side by side, so axis 1 comes first: order=(1, 0). A 1D tensor has only one axis: order=(0,).
order is a layout hint. Triton uses it to pick fast ways to load and store.
.advance() tells Triton how to move the tile
ptr.advance((d_row, d_col)) moves the tile by that many elements along each axis. Only the offsets change. base, shape, strides, block_shape and order stay the same.
Blue is the starting tile. Moving by (0, D_TILE_SIZE) gives orange. Moving by (ROWS_TILE_SIZE, 0) gives purple.
advance does not change the pointer you call it on. It returns a new one. So always save the result: x_block_ptr = x_block_ptr.advance(...).
In weighted_sum_fwd, each loop step moves x one tile to the right with (0, D_TILE_SIZE), and moves weight with (D_TILE_SIZE,). Both move by the same amount, so the columns of x and the entries of weight stay lined up. The downward move (purple) is not in this kernel. Instead, the bands of rows are spread across parallel programs, which is managed through offsets.
# D = 8, D_TILE_SIZE = 4, so the loop runs cdiv(8, 4) = 2 times
i = 0 load x at (0, 0), then advance to (0, 4)
i = 1 load x at (0, 4), then advance to (0, 8)
# loop ends. (0, 8) is past the edge, but it is never loaded.
Moving the tile past the edge is fine; only loading or storing there needs care.
boundary_check tells Triton what happens past the edge
These two arguments belong to tl.load and tl.store, not to make_block_ptr. But they only work because you gave the block pointer a shape. boundary_check lists the axes to check, i.e. 0 means rows and 1 means columns. padding_option says what to put in the spots that are outside, e.g. "zero" for 0, or "nan" for NaN ("not a number", a special float value that turns any sum it touches into NaN).
I filled x with 0, 1, 2, 3, ... so you can see where each loaded number came from.
"zero" is right for this kernel because the outside spots become 0 in both row and weight, so row * weight adds nothing to the sum.
tl.store with boundary_check skips spots outside shape. The kernel stores output with boundary_check=(0,) because the last program may have rows past NUM_ROWS, and those must not be written.
If you check an axis but leave out padding_option, memory past the edge is not read, but the values in those spots are undefined.
Put it all together
This is the whole weighted_sum_fwd kernel. Each program sums a band of rows of x, weighted by weight, one tile at a time. Change the sizes, pick a program, and step through the loop.
Cheat sheet
| Argument | In plain words | x | weight | output |
|---|---|---|---|---|
base |
Address of element 0 | x_ptr |
weight_ptr |
output_ptr |
shape |
Size of the whole tensor | (NUM_ROWS, D) |
(D,) |
(NUM_ROWS,) |
strides |
Memory jump for one step on each axis | (x_stride_row, x_stride_dim) |
(weight_stride_dim,) |
(output_stride_row,) |
offsets |
Where the tile starts, in elements | (row_tile_idx * ROWS_TILE_SIZE, 0) |
(0,) |
(row_tile_idx * ROWS_TILE_SIZE,) |
block_shape |
Tile size, fixed at compile time | (ROWS_TILE_SIZE, D_TILE_SIZE) |
(D_TILE_SIZE,) |
(ROWS_TILE_SIZE,) |
order |
Axes from fastest to slowest in memory | (1, 0) |
(0,) |
(0,) |
Recipe: a 2D row-major tensor
ptr = tl.make_block_ptr(
base_ptr,
shape=(M, N),
strides=(stride_m, stride_n), # t.stride(0), t.stride(1)
offsets=(pid * BM, 0), # in elements
block_shape=(BM, BN), # tl.constexpr, powers of 2
order=(1, 0), # axis 1 has stride 1
)
tile = tl.load(ptr, boundary_check=(0, 1), padding_option="zero")
ptr = ptr.advance((0, BN)) # save the result
Recipe: a 1D tensor
ptr = tl.make_block_ptr(
base_ptr,
shape=(N,), # note the comma: (N) is just N
strides=(stride,), # t.stride(0), usually 1
offsets=(start,), # in elements
block_shape=(B,), # tl.constexpr, a power of 2
order=(0,), # only one axis
)
tile = tl.load(ptr, boundary_check=(0,), padding_option="zero")
ptr = ptr.advance((B,)) # save the result; note the comma again
Common mistakes with block pointers
The examples use a 2D tensor x of shape (N, D), split into tiles of rows. After the first example, snippets show only the arguments that matter.
- Passing the tile's address as
base. Pass the start of the tensor and useoffsetsfor the tile.
row_tile_idx = tl.program_id(0)
# Wrong: base already points at the tile
tile_ptr = x_ptr + row_tile_idx * ROWS_TILE_SIZE * stride_row
x_block_ptr = tl.make_block_ptr(
base=tile_ptr,
shape=(N, D),
strides=(stride_row, stride_col),
offsets=(0, 0),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
)
# Right: base is the tensor start, offsets locate the tile
x_block_ptr = tl.make_block_ptr(
base=x_ptr,
shape=(N, D),
strides=(stride_row, stride_col),
offsets=(row_tile_idx * ROWS_TILE_SIZE, 0),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
)
- Putting a tile number in
offsets. Userow_tile_idx * ROWS_TILE_SIZE, notrow_tile_idx.
# Wrong: tile 3 starts at row 3
offsets=(row_tile_idx, 0),
# Right: tile 3 starts at row 3 * ROWS_TILE_SIZE
offsets=(row_tile_idx * ROWS_TILE_SIZE, 0),
- Putting the tile size in
shape.shapeis the whole tensor. The tile size goes inblock_shape.
# Wrong: shape describes the tile
shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
# Right: shape describes the whole tensor
shape=(N, D),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
- Using a runtime value in
block_shape. It must betl.constexpr, and in practice a power of 2.
# Wrong: tile sizes are runtime arguments
@triton.jit
def kernel(x_ptr, N, D, stride_row, stride_col, ROWS_TILE_SIZE, D_TILE_SIZE):
...
# Right: tile sizes are constexpr
@triton.jit
def kernel(x_ptr, N, D, stride_row, stride_col,
ROWS_TILE_SIZE: tl.constexpr, D_TILE_SIZE: tl.constexpr):
...
# On the host, round up to a power of 2
D_TILE_SIZE = triton.next_power_of_2(D)
- Writing
(D)for a 1D tuple. Python reads that as a plain number. Write(D,).
# Wrong: every "tuple" here is a plain number
shape=(D),
strides=(stride),
offsets=(tile_idx * TILE_SIZE),
block_shape=(TILE_SIZE),
order=(0),
# Right
shape=(D,),
strides=(stride,),
offsets=(tile_idx * TILE_SIZE,),
block_shape=(TILE_SIZE,),
order=(0,),
- Giving the tuples different lengths.
shape,strides,offsets,block_shape, andorderneed one entry per dimension.
# Wrong: strides and offsets have one entry for a 2D tensor
shape=(N, D),
strides=(stride_row,),
offsets=(row_tile_idx * ROWS_TILE_SIZE,),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
# Right: two entries everywhere
shape=(N, D),
strides=(stride_row, stride_col),
offsets=(row_tile_idx * ROWS_TILE_SIZE, 0),
block_shape=(ROWS_TILE_SIZE, D_TILE_SIZE),
order=(1, 0),
- Assuming
stride_row = D. Pass the real strides fromx.stride().
x = y.t() # shape (N, D), but x.stride() == (1, N)
# Wrong: assumes x is contiguous
kernel[grid](x, N, D, D, 1, ROWS_TILE_SIZE=32, D_TILE_SIZE=64)
# Right: works for transposed and sliced tensors too
kernel[grid](x, N, D, x.stride(0), x.stride(1), ROWS_TILE_SIZE=32, D_TILE_SIZE=64)
- Calling
ptr.advance(...)without saving the result. It returns a new pointer, so writeptr = ptr.advance(...).
# Wrong: the result is thrown away, so every iteration loads the first tile
for _ in range(tl.cdiv(D, D_TILE_SIZE)):
x = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
x_block_ptr.advance((0, D_TILE_SIZE))
# Right
for _ in range(tl.cdiv(D, D_TILE_SIZE)):
x = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
x_block_ptr = x_block_ptr.advance((0, D_TILE_SIZE))
- Putting tile counts in
advance. Its offsets are in elements too. Useptr.advance((0, D_TILE_SIZE)), notptr.advance((0, 1)).
# Wrong: moves one column
x_block_ptr = x_block_ptr.advance((0, 1))
# Right: moves one tile
x_block_ptr = x_block_ptr.advance((0, D_TILE_SIZE))
- Skipping
boundary_checkwhen the sizes do not divide evenly.
# Wrong: no check, or booleans instead of dimension indices
x = tl.load(x_block_ptr)
x = tl.load(x_block_ptr, boundary_check=(True, True))
tl.store(out_block_ptr, y)
# Right: check dims 0 and 1 on both the load and the store
x = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
tl.store(out_block_ptr, y, boundary_check=(0, 1))
- Using
padding_option="nan"for a sum. One NaN turns the whole sum into NaN.
# Wrong: padding is undefined, or NaN
x = tl.load(x_block_ptr, boundary_check=(0, 1))
x = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="nan")
acc += tl.sum(x, axis=1)
# Right
x = tl.load(x_block_ptr, boundary_check=(0, 1), padding_option="zero")
acc += tl.sum(x, axis=1)
Step 2: the Triton forward kernel
The Triton version has the same shape as the PyTorch one, with the outer loop over query tiles replaced by the launch grid. Each program (Triton's name for a thread block) gets one query tile and one batch element, loads Q_i once, and loops over key tiles on its own.
Below is the kernel as it stands at the end of this article; the two lines marked FIX will be explained in Step 3.
# flashattention_autograd_function_triton.py
import math
import torch
import triton
import triton.language as tl
@triton.jit
def flash_fwd_kernel(
Q_ptr, K_ptr, V_ptr, O_ptr, L_ptr,
stride_qb, stride_qq, stride_qd,
stride_kb, stride_kk, stride_kd,
stride_vb, stride_vk, stride_vd,
stride_ob, stride_oq, stride_od,
stride_lb, stride_lq,
N_QUERIES, N_KEYS,
scale,
D: tl.constexpr,
Q_TILE_SIZE: tl.constexpr,
K_TILE_SIZE: tl.constexpr,
is_causal: tl.constexpr,
):
query_tile_index = tl.program_id(0)
batch_index = tl.program_id(1)
# Block pointers: a (rows, D) tile into each tensor for this batch element.
# Q and O tiles start at this program's query tile; K and V start at row 0
# and are advanced inside the loop.
Q_block_ptr = tl.make_block_ptr(
Q_ptr + batch_index * stride_qb,
shape=(N_QUERIES, D), strides=(stride_qq, stride_qd),
offsets=(query_tile_index * Q_TILE_SIZE, 0),
block_shape=(Q_TILE_SIZE, D), order=(1, 0),
)
K_block_ptr = tl.make_block_ptr(
K_ptr + batch_index * stride_kb,
shape=(N_KEYS, D), strides=(stride_kk, stride_kd),
offsets=(0, 0), block_shape=(K_TILE_SIZE, D), order=(1, 0),
)
V_block_ptr = tl.make_block_ptr(
V_ptr + batch_index * stride_vb,
shape=(N_KEYS, D), strides=(stride_vk, stride_vd),
offsets=(0, 0), block_shape=(K_TILE_SIZE, D), order=(1, 0),
)
O_block_ptr = tl.make_block_ptr(
O_ptr + batch_index * stride_ob,
shape=(N_QUERIES, D), strides=(stride_oq, stride_od),
offsets=(query_tile_index * Q_TILE_SIZE, 0),
block_shape=(Q_TILE_SIZE, D), order=(1, 0),
)
L_block_ptr = tl.make_block_ptr(
L_ptr + batch_index * stride_lb,
shape=(N_QUERIES,), strides=(stride_lq,),
offsets=(query_tile_index * Q_TILE_SIZE,),
block_shape=(Q_TILE_SIZE,), order=(0,),
)
# Running state, kept in fp32 regardless of input dtype.
O_acc = tl.zeros((Q_TILE_SIZE, D), dtype=tl.float32)
l_acc = tl.zeros((Q_TILE_SIZE, 1), dtype=tl.float32)
m_acc = tl.full((Q_TILE_SIZE, 1), value=float("-inf"), dtype=tl.float32)
Q_i = tl.load(Q_block_ptr, boundary_check=(0, 1), padding_option="zero")
q_pos = (query_tile_index * Q_TILE_SIZE + tl.arange(0, Q_TILE_SIZE))[:, None]
for j in range(tl.cdiv(N_KEYS, K_TILE_SIZE)):
k_pos = (j * K_TILE_SIZE + tl.arange(0, K_TILE_SIZE))[None, :]
K_j = tl.load(K_block_ptr, boundary_check=(0, 1), padding_option="zero")
V_j = tl.load(V_block_ptr, boundary_check=(0, 1), padding_option="zero")
S_ij = tl.dot(Q_i, tl.trans(K_j)) * scale # (Q_TILE, K_TILE), fp32
# FIX 2: zero-padded keys past N_KEYS score 0, not -inf. Mask them.
keep = k_pos < N_KEYS
if is_causal:
keep = keep & (k_pos <= q_pos)
S_ij = tl.where(keep, S_ij, -1e6)
m_new = tl.maximum(m_acc, tl.max(S_ij, axis=1, keep_dims=True))
P_ij = tl.exp(S_ij - m_new)
alpha = tl.exp(m_acc - m_new)
l_acc = alpha * l_acc + tl.sum(P_ij, axis=1, keep_dims=True)
# FIX 1: the cast must be assigned. tl.dot needs both operands in the
# same dtype; the fp32 accumulator is passed separately via acc=.
P_ij = P_ij.to(V_j.dtype)
O_acc = alpha * O_acc
O_acc = tl.dot(P_ij, V_j, acc=O_acc)
m_acc = m_new
K_block_ptr = K_block_ptr.advance((K_TILE_SIZE, 0))
V_block_ptr = V_block_ptr.advance((K_TILE_SIZE, 0))
O_i = (O_acc / l_acc).to(O_block_ptr.type.element_ty)
tl.store(O_block_ptr, O_i, boundary_check=(0, 1))
L_i = tl.reshape(m_acc + tl.log(l_acc), (Q_TILE_SIZE,))
tl.store(L_block_ptr, L_i, boundary_check=(0,))
class FlashAttentionTriton(torch.autograd.Function):
Q_TILE_SIZE = 16
K_TILE_SIZE = 16
@staticmethod
def forward(ctx, Q, K, V, is_causal=False):
assert Q.ndim == 3, "expects (batch, seq, head_dim); flatten (B, H, N, D) to (B*H, N, D)"
assert Q.stride(-1) == 1 and K.stride(-1) == 1 and V.stride(-1) == 1
B, N_q, D = Q.shape
N_k = K.shape[1]
assert D in (16, 32, 64, 128), "block_shape dims must be powers of two"
O = torch.empty_like(Q)
L = torch.empty((B, N_q), device=Q.device, dtype=torch.float32)
grid = (triton.cdiv(N_q, FlashAttentionTriton.Q_TILE_SIZE), B)
flash_fwd_kernel[grid](
Q, K, V, O, L,
Q.stride(0), Q.stride(1), Q.stride(2),
K.stride(0), K.stride(1), K.stride(2),
V.stride(0), V.stride(1), V.stride(2),
O.stride(0), O.stride(1), O.stride(2),
L.stride(0), L.stride(1),
N_q, N_k,
1.0 / math.sqrt(D),
D=D,
Q_TILE_SIZE=FlashAttentionTriton.Q_TILE_SIZE,
K_TILE_SIZE=FlashAttentionTriton.K_TILE_SIZE,
is_causal=is_causal,
)
ctx.save_for_backward(Q, K, V, O, L)
ctx.is_causal = is_causal
return O
@staticmethod
def backward(ctx, dO):
raise NotImplementedError("tiled backward is left for future articles")
Legend:
| Term | Meaning |
|---|---|
| grid | How many programs to launch along each axis. Here, (cdiv(N_q, 16), B*H). |
program_id(0) |
Which query tile this program handles: 0, 1, 2, … |
program_id(1) |
Which batch element this program handles. After flattening, that means one (sequence, head) pair. |
| query tile | A block of 16 consecutive query rows (tokens), handled by one program. |
| 16 | The query tile size: rows of Q per program. |
cdiv(a, b) |
a ÷ b, rounded up, so a partly filled last tile still gets a program. |
B |
Batch size: the number of sequences processed together. |
H |
The number of attention heads. |
N |
Sequence length: tokens per sequence. |
N_q |
The number of query tokens. In self-attention it equals N. |
D |
Head dimension: how many numbers each token has in one head. |
(B, H, N, D) |
The shape of Q, K and V: batch, heads, tokens, head dimension. |
(B*H, N, D) |
Batch and heads merged into one axis, so each (sequence, head) pair is treated as a separate batch element. |
Reading it top to bottom:
The grid. program_id(0) indexes query tiles, program_id(1) indexes batch elements, so the launch has cdiv(N_q, 16) × B programs. Multi-head attention is handled by flattening (B, H, N, D) into (B*H, N, D) before the call. With the flattening, the kernel never knows about heads. At 16 rows per program and 4096 tokens × 32 heads, that is 8192 programs, which is plenty to fill a GPU. The problem, as I show in the last section, is that a 16x16 tile is too small to saturate the GPU.
Block pointers. make_block_ptr describes a 2D tile of a strided tensor, i.e. the base pointer, the full logical shape, the strides, where the tile starts, and how big it is. order=(1, 0) says the last dimension is the fastest-varying one, i.e. row-major. The payoff is boundary_check and padding_option on load, which handle the ragged last tile for free, and .advance(), which slides the K and V tiles down by one tile per iteration without recomputing addresses.
Accumulators in fp32. O_acc, l_acc and m_acc are fp32 whatever the input dtype, and tl.dot returns fp32 for bf16 inputs. This is what makes it safe to run in bf16. The inputs and the P·V product are lower precision, but the sum across key tiles is higher precision.
One Q load, many K/V loads. Q_i is loaded once before the loop. K_j and V_j are loaded per iteration, and every program loads the entire K and V for its batch element. That is inherent to the algorithm: it is the trade-off FlashAttention-2 makes to avoid the N×N matrix, which is why K/V loads dominate the memory traffic and why tile size matters.
The inner loop is a line-for-line translation of the PyTorch version: score block, mask, new max, unnormalised probabilities, rescale factor, update l, update O, carry m forward. tl.dot(P_ij, V_j, acc=O_acc) fuses the matrix multiply and the accumulate.
The store. boundary_check on the store discards rows of the last tile that fall past N_QUERIES, so the padded rows computed with zeroed Q are never written.
Step 3: testing and bugs found
The testing harness ran one shape, (4, 128, 128, 64), in fp32, at a tolerance of 1e-2, and my kernel passed the test. However, the kernel was still wrong in two ways, and Claude found both bugs using the test file below, which was also generated by Claude.
# test_flashattention_triton.py
import math
import pytest
import torch
import torch.nn.functional as F
from flashattention_autograd_function_triton import FlashAttentionTriton
def reference(q, k, v, is_causal):
o = F.scaled_dot_product_attention(q, k, v, is_causal=is_causal)
s = (q.float() @ k.float().transpose(-1, -2)) / math.sqrt(q.shape[-1])
if is_causal:
tril = torch.ones(s.shape[-2:], dtype=torch.bool, device=q.device).tril()
s = s.masked_fill(~tril, float("-inf"))
return o, torch.logsumexp(s, dim=-1)
# (batch, n_queries, n_keys, head_dim). Deliberately includes shapes that
# are not multiples of the 16-row tile and non-square attention.
SHAPES = [
(4, 128, 128, 64), # the assignment's shape
(2, 1024, 1024, 64),
(2, 100, 100, 64), # ragged last tile on both axes
(1, 37, 53, 32), # ragged and non-square
(1, 200, 96, 128),
]
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
@pytest.mark.parametrize("is_causal", [False, True])
@pytest.mark.parametrize("B,Nq,Nk,D", SHAPES)
def test_forward_matches_sdpa(B, Nq, Nk, D, is_causal, dtype):
if is_causal and Nq != Nk:
pytest.skip("keep causal tests square so the mask convention is unambiguous")
torch.manual_seed(0)
q = torch.randn(B, Nq, D, device="cuda", dtype=dtype, requires_grad=True)
k = torch.randn(B, Nk, D, device="cuda", dtype=dtype, requires_grad=True)
v = torch.randn(B, Nk, D, device="cuda", dtype=dtype, requires_grad=True)
o = FlashAttentionTriton.apply(q, k, v, is_causal)
# L is not returned; pull it out of the saved tensors, as the assignment test does.
L = [t for t in o.grad_fn.saved_tensors if t.shape == (B, Nq)][0]
o_ref, L_ref = reference(q, k, v, is_causal)
# fp32 tl.dot uses TF32 on tensor cores by default; bf16 rounds P before P@V.
tol = dict(atol=2e-2, rtol=2e-2) if dtype == torch.bfloat16 else dict(atol=1e-2, rtol=1e-2)
torch.testing.assert_close(o, o_ref, **tol)
torch.testing.assert_close(L, L_ref, **tol)
Two details in the harness:
-
Lis not part of the function's return value, so the test digs it out ofo.grad_fn.saved_tensors, which is the same trick the assignment's test uses and requires the inputs to haverequires_grad=True. - The fp32 tolerance is looser than you might expect because Triton's
tl.doton fp32 inputs uses TF32 on Ampere and later unless you passinput_precision="ieee".
Bug 1: the casts that never happened
In my original code which passed the harness test, I had the following three lines:
P_ij.to(V_j.type.element_ty) # "cast P to bf16 before the dot"
...
O_i.to(O_block_ptr.type.element_ty)
...
L_i.to(L_block_ptr.type.element_ty)
In Triton, as in PyTorch, .to() returns a new tensor. .to() does not mutate. All three lines were no-ops, and the comment on the first line described a cast that was not taking place. The kernel still passed the harness test because the test ran in fp32, where casting fp32 to fp32 changed nothing.
The bf16 test found it at compile time rather than as a numerical error. tl.dot required both operands to share a dtype, and after the no-op P_ij was still fp32 while V_j was bf16. The fix: P_ij = P_ij.to(V_j.dtype). The O and L casts are fixed the same way, folded into the store lines.
The lesson I took from this is to check the dtype of every tl.dot operand when a kernel is first written.
Bug 2: padded keys that voted
tl.load with boundary_check and padding_option="zero" fills rows past the end of the tensor with zeros. For Q that is harmless because padded query rows produce garbage that the store then discards. For K it is problematic because a zero key row gives a score of exactly 0 for every query, so the padded keys enter the softmax with weight exp(0 - m) each. When N_KEYS isn't a multiple of the key tile (say 100 keys in tiles of 16), the last tile runs past the real keys, and its extra slots are "phantom" keys that do not exist. Unless they are masked out, the output becomes a weighted average over real keys plus these phantom ones.
The original harness test did not catch this because 128 is a multiple of 16. However, on a (2, 100, 100, 64) tensor, the kernel failed. The fix is the keep = k_pos < N_KEYS mask, applied whether or not attention is causal. It costs one compare and one select per score, which is nothing next to the two matrix multiplies, and it means the kernel is now correct for any sequence length, not just multiples of 16.
Step 4: the benchmark, and the baseline number
This section reports the kernel's speed as a percentage of PyTorch's scaled_dot_product_attention on the same inputs and the same NVIDIA RTX 4070 Super. SDPA is chosen because at head dim 64 in bf16 it dispatches to the FlashAttention-2 kernel that ships inside PyTorch, so "100% of SDPA" means "as fast as FA2". The harness is short and it never changes after this article.
# bench_flashattention.py
import torch
import torch.nn.functional as F
import triton
from flashattention_autograd_function_triton import FlashAttentionTriton
def attn_flops(batch, n, d, is_causal):
# QK^T and PV: two matmuls of n*n*d multiply-adds each, 2 flops per MAC.
f = 4 * batch * n * n * d
return f / 2 if is_causal else f
def bench(B, H, N, D, is_causal, dtype=torch.bfloat16):
q, k, v = (torch.randn(B * H, N, D, device="cuda", dtype=dtype) for _ in range(3))
ms_ours = triton.testing.do_bench(
lambda: FlashAttentionTriton.apply(q, k, v, is_causal), return_mode="median")
q4, k4, v4 = (t.view(B, H, N, D) for t in (q, k, v))
ms_sdpa = triton.testing.do_bench(
lambda: F.scaled_dot_product_attention(q4, k4, v4, is_causal=is_causal),
return_mode="median")
flops = attn_flops(B * H, N, D, is_causal)
tflops = lambda ms: flops / ms * 1e-9
return ms_ours, ms_sdpa, tflops(ms_ours), tflops(ms_sdpa)
if __name__ == "__main__":
print(f"{'N':>6} {'causal':>7} {'ours ms':>9} {'sdpa ms':>9} {'ours TF/s':>10} {'sdpa TF/s':>10} {'% sdpa':>7}")
for N in (512, 1024, 2048, 4096):
for causal in (False, True):
a, b, c, d = bench(B=4, H=8, N=N, D=64, is_causal=causal)
print(f"{N:>6} {str(causal):>7} {a:>9.3f} {b:>9.3f} {c:>10.1f} {d:>10.1f} {100 * b / a:>6.0f}%")
triton.testing.do_bench handles warm-up and flushes L2 between runs, and with return_mode="median" it reports the median, so a single call is enough. The FLOP count is the standard one for attention: two matmuls of N×N×d multiply-adds, halved for causal because the kernel is only supposed to do the lower triangle. Note the word supposed. This kernel computes every tile and masks half of them, so its causal FLOP rate is flattered by the accounting.
Results on the RTX 4070 Super, batch 4, 8 heads, head dim 64, bf16:
| N | causal | mine (ms) | SDPA (ms) | mine TFLOP/s | SDPA TFLOP/s | % of SDPA |
|---|---|---|---|---|---|---|
| 512 | no | 0.107 | 0.063 | 20.1 | 34.3 | 59% |
| 512 | yes | 0.106 | 0.055 | 10.1 | 19.4 | 52% |
| 1024 | no | 0.386 | 0.152 | 22.3 | 56.5 | 39% |
| 1024 | yes | 0.385 | 0.121 | 11.2 | 35.5 | 31% |
| 2048 | no | 1.534 | 0.539 | 22.4 | 63.7 | 35% |
| 2048 | yes | 1.560 | 0.329 | 11.0 | 52.2 | 21% |
| 4096 | no | 6.195 | 2.009 | 22.2 | 68.4 | 32% |
| 4096 | yes | 6.155 | 1.120 | 11.2 | 61.3 | 18% |
At N = 4096 without causal masking, the kernel runs at 32% of SDPA (22.2 vs 68.4 TFLOP/s). The gap widens as N grows: my kernel reaches about 22 TFLOP/s by N = 1024 and stays there, while SDPA keeps climbing from 34 to 68 TFLOP/s, so the ratio falls from 59% at N = 512 to 32% at N = 4096. Causal attention is worse again, at 18% of SDPA at N = 4096. My kernel's causal runtime is the same as its non-causal runtime at every N (6.16 ms vs 6.20 ms at N = 4096), because every tile above the diagonal is still computed and then masked, while SDPA's causal runtime is roughly half its non-causal runtime.
To check which SDPA backend you're actually comparing against, wrap the reference call in torch.nn.attention.sdpa_kernel([SDPBackend.FLASH_ATTENTION]): if it errors, you're not comparing against FA2 and the percentages mean something else.
What is wrong with this kernel
The kernel is correct and it is slow. Here is why, in the order I plan to fix it. Each item names the cause, not just the symptom, so that future articles can test it.
1. The tiles are too small. 16×16 tiles mean each tl.dot is a 16×64 by 64×16 product. Tensor cores want bigger operands than that to reach their throughput, and each program does so little work per K/V load that the kernel spends its time moving data rather than multiplying it. Every program also reloads all of K and V from L2 or HBM, and with a 16-row Q tile the ratio of loads to useful FLOPs is about eight times worse than with a 128-row tile. The right sizes depend on the card: Ada has roughly 100 KB of shared memory per SM against Hopper's 228 KB, so the 128×64 or 128×128 configurations in the FA2 paper are not automatically right here. In future articles I will dive into the shared-memory budget arithmetic, sweep tile sizes together with num_warps and num_stages, and wrap the winner in @triton.autotune.
2. Causal attention does all the work and throws half away. The mask is applied to every tile, including tiles that lie entirely above the diagonal, where every element is masked. Those tiles contribute nothing to the output and the kernel still loads K and V for them, runs both matmuls and the exponentials, then discards it all. The fix is to stop the key loop at the diagonal, which roughly halves causal runtime, and to only apply the mask on the one tile that actually straddles it.
3. exp instead of exp2. The GPU's fast exponential unit computes 2^x; tl.exp is exp2(x · log₂e), an extra multiply on every score. The standard trick is to fold scale · log₂e into the scaling of Q once per tile, then use tl.math.exp2 directly. It's a small win but it's in the inner loop.
4. The -1e6 mask sentinel. Masking with a large negative number rather than -inf sidesteps a real hazard: a row whose every score is -inf produces exp(-inf - (-inf)) = NaN. A fully masked tile is harmless once a row has seen a real key, because exp(-1e6 - m) underflows to exactly zero in fp32. The remaining risk is a row with no valid key at all: every score is then -1e6, the row maximum is -1e6, and each masked key gets weight exp(0) = 1, so the row silently becomes an average of values it should never see. In this kernel every row can see at least key 0, so it cannot happen yet, but kernels that allow such rows need to handle them explicitly.
5. No backward pass. backward raises NotImplementedError, and the PyTorch reference backward materialises the full N×N score matrix, so it isn't FlashAttention's backward at all. The tiled backward needs the L I have been carrying, a precomputed D = rowsum(dO ∘ O), and a kernel that iterates over query tiles for each key tile to build dK and dV, with dQ accumulated either by atomics or by a second pass.
6. The mask is computed unconditionally. The k_pos < N_KEYS compare and the tl.where run on every tile, including the full interior ones where nothing is masked. It is cheap, but in a kernel this small everything in the inner loop counts. Once tile-skipping is in place the mask can be restricted to boundary tiles.
7. No profile. Everything above is reasoning from first principles. None of it has been confirmed with Nsight Compute, and reasoning about GPU performance without a profiler is how people end up optimising the wrong thing. In future articles I will profile the forward and backward passes, compare these kernels against the GPU's roofline, and identify what is left between this kernel and FA2.
Future articles
Some ideas I have for future articles include:
- Take the kernel above and change the tile sizes,
num_warpsandnum_stages. Work through the shared-memory budget for an Ada SM, sweep the configuration space, show the heatmap, then wrap the winner in@triton.autotuneand measure what that costs in compile time. The benchmark harness and the test file stay exactly as they are here, so the numbers are directly comparable to the ones above. - Go inside the inner loop for causal tile-skipping,
exp2and the P cast. - Backward pass of FA2.
- Profiler and the roofline.
- Rewrite the kernel in TileLang and CuTe DSL.







Top comments (0)