DEV Community

Cover image for FlashAttention 2 from PyTorch to Triton
Lewis Won
Lewis Won

Posted on AI-assisted

FlashAttention 2 from PyTorch to Triton

Table of Contents

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.

FlashAttention-2 Forward Pass

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
Enter fullscreen mode Exit fullscreen mode

A few points to note:

  • m_i starts at negative infinity so that the first tile's row maximum "wins" in line 10 of the algorithm above. m_i starting at negative infinity also ensures on the first step alpha = 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_i is the weighted sum of value rows, and the softmax needs that sum divided by the total weight l_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 a B_q x d tile on every inner loop.
  • L is stored, not m and l separately. The backward pass recomputes each tile of S from Q and K, then recovers that tile's probabilities as exp(S - L). For context, assuming N = 4096, S with dimension NxN per head per batch element will have about 16.8 million entries, or about 33MB in bf16. L with dimension Nx1 per head per batch element contains 4096 entries, or about 16KB.
  • The is_causal flag 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
Enter fullscreen mode Exit fullscreen mode

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.

01-overview

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)
Enter fullscreen mode Exit fullscreen mode

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,
)
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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)):
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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]
Enter fullscreen mode Exit fullscreen mode

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]]
Enter fullscreen mode Exit fullscreen mode
  • cols[None, :] lays the list down as a row (1 row, 4 columns):
[[8, 9, 10, 11]]
Enter fullscreen mode Exit fullscreen mode

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]]
Enter fullscreen mode Exit fullscreen mode

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
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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]]
Enter fullscreen mode Exit fullscreen mode

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.

6-without-block-ptr

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)
Enter fullscreen mode Exit fullscreen mode

This gives [weight[8], weight[9], 0, 0].

output += tl.sum(row * weight[None, :], axis=1)
Enter fullscreen mode Exit fullscreen mode

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)
Enter fullscreen mode Exit fullscreen mode

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 an other value 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,))
Enter fullscreen mode Exit fullscreen mode

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.

2-base-memory

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.


Pass the start of the whole tensor, not the start of your tile. You choose the tile later, with offsets.

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,).

3-shape

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.

address of x[r,c]=base+r×row stride+c×column stride \text{address of } x[r, c] = \text{base} + r \times \text{row stride} + c \times \text{column stride}

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).


Do not assume stride_row = D. Transposed and sliced tensors break that rule. With the real strides, the kernel reads the right numbers.

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).

4-order

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.

5-advance

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.
Enter fullscreen mode Exit fullscreen mode

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
Enter fullscreen mode Exit fullscreen mode

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
Enter fullscreen mode Exit fullscreen mode

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 use offsets for 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),
  )
Enter fullscreen mode Exit fullscreen mode
  • Putting a tile number in offsets. Use row_tile_idx * ROWS_TILE_SIZE, not row_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),
Enter fullscreen mode Exit fullscreen mode
  • Putting the tile size in shape. shape is the whole tensor. The tile size goes in block_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),
Enter fullscreen mode Exit fullscreen mode
  • Using a runtime value in block_shape. It must be tl.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)
Enter fullscreen mode Exit fullscreen mode
  • 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,),
Enter fullscreen mode Exit fullscreen mode
  • Giving the tuples different lengths. shape, strides, offsets, block_shape, and order need 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),
Enter fullscreen mode Exit fullscreen mode
  • Assuming stride_row = D. Pass the real strides from x.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)
Enter fullscreen mode Exit fullscreen mode
  • Calling ptr.advance(...) without saving the result. It returns a new pointer, so write ptr = 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))
Enter fullscreen mode Exit fullscreen mode
  • Putting tile counts in advance. Its offsets are in elements too. Use ptr.advance((0, D_TILE_SIZE)), not ptr.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))
Enter fullscreen mode Exit fullscreen mode
  • Skipping boundary_check when 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))
Enter fullscreen mode Exit fullscreen mode
  • 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)
Enter fullscreen mode Exit fullscreen mode

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")
Enter fullscreen mode Exit fullscreen mode

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.


GPUs usually start programs in order with axis 0 changing fastest, so putting query tiles on axis 0 makes programs that run together share the same K and V (better cache reuse), and on NVIDIA GPUs axis 0 is also the only one allowed more than 65,535 programs, which long sequences can need.

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)
Enter fullscreen mode Exit fullscreen mode

Two details in the harness:

  1. L is not part of the function's return value, so the test digs it out of o.grad_fn.saved_tensors, which is the same trick the assignment's test uses and requires the inputs to have requires_grad=True.
  2. The fp32 tolerance is looser than you might expect because Triton's tl.dot on fp32 inputs uses TF32 on Ampere and later unless you pass input_precision="ieee".

TF32 keeps fp32's range but rounds each input to about 3 significant digits before multiplying, which makes matrix maths much faster on Tensor Cores but leaves errors around 1e-3. input_precision="ieee" forces true fp32 maths, which is slower but matches the reference, i.e. the plain PyTorch version the test compares against runs in true fp32.

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)
Enter fullscreen mode Exit fullscreen mode

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}%")
Enter fullscreen mode Exit fullscreen mode

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_warps and num_stages. Work through the shared-memory budget for an Ada SM, sweep the configuration space, show the heatmap, then wrap the winner in @triton.autotune and 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, exp2 and the P cast.
  • Backward pass of FA2.
  • Profiler and the roofline.
  • Rewrite the kernel in TileLang and CuTe DSL.

Top comments (0)