2Day 2

SDPA and FlashAttention's tiling

Plain attention builds a T by T score matrix in memory. FlashAttention computes the same result a block at a time with a running maximum and sum, and never stores the square. Day 2 wrote the tiling by hand and measured PyTorch's kernels against the plain path.

Causal attention, by hand computed the full score matrix, masked it, took a softmax, and multiplied by the values. That is correct and easy to read. It also allocates a T×TT \times T matrix per head, and at long lengths that matrix dominates both memory and time. Day 2 replaced it with PyTorch's fused attention function and kept the handwritten path as the reference.

The problem

Attention is limited by memory traffic, not arithmetic

The FlashAttention paper's argument starts with the GPU's memory layout. An A100 has 40 to 80 GB of main memory (HBM) at 1.5 to 2.0 TB/s, and 192 KB of fast on-chip memory (SRAM) per multiprocessor at about 19 TB/s. Plain attention writes the score matrix to HBM, reads it back for the softmax, writes the probabilities, and reads them again to multiply by VV. For long sequences those reads and writes, not the multiplications, set the time. The paper counts Θ(Nd+N2)\Theta(Nd + N^2) memory accesses for the standard method.

The fix is to never write the square. Load a block of keys and values into fast memory, compute that block's scores for a block of queries, fold them into the output, and move on. The obstacle is softmax, which needs the maximum and the sum over the whole row before it can normalize anything.

Playground
HBM, large and slow16 GB on a T4Q, K, V, output, and S, P if standardon-chip SRAM, small, fastone 64-row tile at a timereadwrite
standard17.3M
tiled8.7M
standard, elements moved17.3MQ, K, V, O plus S and P written and read
tiled, elements moved8.7MK and V reread once per query block
Ratio2.0×standard ÷ tiled

Element counts for one head of width 64, simplified from the FlashAttention paper's analysis. Standard attention writes the T × T score matrix to HBM, reads it back for the softmax, writes the probabilities and reads them again. The tiled version never stores the square. It pays by rereading K and V once per query block, which is why larger tiles, if SRAM fits them, move less.

Where attention's data lives on a GPU, and how many elements each method moves between the two memories.
Online softmax

A softmax you can compute one block at a time

Softmax subtracts the row maximum before exponentiating, for numerical safety. Suppose we have processed some keys and kept three things per query row: the largest score so far mm, the sum ℓ=∑esj−m\ell = \sum e^{s_j - m}, and the unnormalized output o=∑esj−mvjo = \sum e^{s_j - m} v_j. A new block of scores arrives with maximum m~\tilde m. Then:

m′=max⁡(m,m~),ℓ′=em−m′ ℓ+∑j∈blockesj−m′,o′=em−m′ o+∑j∈blockesj−m′ vjm' = \max(m, \tilde m), \qquad \ell' = e^{m - m'}\,\ell + \sum_{j \in \text{block}} e^{s_j - m'}, \qquad o' = e^{m - m'}\,o + \sum_{j \in \text{block}} e^{s_j - m'}\, v_j

The factor em−m′e^{m - m'} rescales everything computed under the old maximum to the new one. After the last block, o/ℓo / \ell is exactly the softmax-weighted average. No step needed the whole row.

Playground
Block size
k0−1.19
k10.95
k2−0.27
k3−0.12
k40.35
k51.07
k6−1.26
k70.37
k8−1.39
k9−1.25
k10−0.49
k11−1.39
Keys read so far0 to 3block 1 of 3
Running max m0.953largest score seen
Running sum ℓ1.754Σ exp(score − m)
Output so far, entry 00.6495full softmax: 0.1942
Each block rescales the old sum and output by exp(m_old − m_new).
One query row of a 12-position causal attention, processed in blocks. The strip shows that row's scores: gray for blocks already folded in, blue for the current block, hatched for masked future keys. The inputs come from the parity fixture, and a test checks this page's tiled result against octlm's tiled_attention.

What happened?

In octlm

The tiled sketch, in Python

Day 2 wrote this forward pass in plain PyTorch to learn the blocking. It is a sketch of the algorithm, not a fast kernel:

octlm/day2.pyline 100
def tiled_attention(query: Tensor, key: Tensor, value: Tensor, block: int = 128) -> Tensor:
    """FlashAttention's blocking. One key block at a time, rescaling by a running maximum."""
    length, head_size = query.shape[-2], query.shape[-1]
    rows = torch.arange(length, device=query.device).unsqueeze(1)
    output = torch.zeros_like(query)
    running_max = torch.full(query.shape[:-1] + (1,), float("-inf"), device=query.device)
    running_sum = torch.zeros_like(running_max)
    for start in range(0, length, block):
        stop = min(start + block, length)
        scores = query @ key[..., start:stop, :].transpose(-2, -1) / math.sqrt(head_size)
        columns = torch.arange(start, stop, device=query.device).unsqueeze(0)
        scores = scores.masked_fill(columns > rows, float("-inf"))
        block_max = torch.maximum(running_max, scores.amax(dim=-1, keepdim=True))
        block_max = torch.where(block_max.isinf(), torch.zeros_like(block_max), block_max)
        correction = (running_max - block_max).exp()
        weights = (scores - block_max).exp()
        running_sum = correction * running_sum + weights.sum(dim=-1, keepdim=True)
        output = correction * output + weights @ value[..., start:stop, :]
        running_max = block_max
    return output / running_sum

It matches scaled_dot_product_attention to 4.8e-7 at lengths 128, 512 and 1,024 and block sizes 64, 128 and 256, from runs/day2-exp014-tiled.jsonl.

The model itself calls PyTorch's fused function. The note read its signature from the installed build rather than trust the documentation: (query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None, enable_gqa=False). octlm passes is_causal=True instead of a mask tensor, never both, and enable_gqa when there are fewer KV heads than query heads.

octlm/model.pyline 312
def forward(self, x: Tensor, rope: tuple[Tensor, Tensor] | None = None) -> Tensor:
    batch, length, width = x.shape
    query, key, value = self._project(x, rope)
    mask = self._mask_for(length, x.device)
    if self.config.attention == "naive":
        attended = self._naive(
            query,
            key,
            value,
            mask if mask is not None else attention_mask(length, 0, 0, x.device),
        )
    else:
        attended = F.scaled_dot_product_attention(
            query,
            key,
            value,
            attn_mask=mask,
            is_causal=mask is None,
            dropout_p=self.dropout.p if self.training else 0.0,
            enable_gqa=key.shape[1] != self.n_heads,
        )
    attended = attended.transpose(1, 2).contiguous().view(batch, length, width)
    return self.output(attended)

Before any Day 2 run used it, the full model through SDPA was checked against the handwritten path at 8, 4, 2 and 1 KV heads. The largest difference was 1.1e-6, from runs/day2-exp014-equivalence.jsonl. That check is what licensed attention = "sdpa" for every later run.

EXP-014

Math against flash on a CPU

PyTorch has several SDPA backends. MATH is the plain algorithm, score matrix and all. FLASH_ATTENTION is the tiled one. On the laptop's CPU build, both ran in float32 and bfloat16, and the memory-efficient backend had no CPU kernel. The benchmark ran one forward pass at batch 1, 8 heads, head width 32, and ran each measurement in its own process: ru_maxrss reports a process's lifetime peak, and the first version of the benchmark let the math backend's peak leak into the flash row.

Measured
mathflash
64 MiB256 MiB1,024 MiB4,096 MiB1,0242,0484,0968,192sequence lengthmemory growthmathflash
Hover the chart to read values.
Resident memory growth during one attention forward pass, float32, CPU, from runs/day2-exp014-sdpa.jsonl. Both axes are logarithmic, base 2. Doubling the length should quadruple a quadratic cost.
Measured
math fp32flash fp32math bf16flash bf16
0.0160.0630.2501.0001,0242,0484,0968,192sequence lengthsecondsmath fp32flash fp32math bf16flash bf16
Hover the chart to read values.
Seconds for one forward pass, float32 and bfloat16, same run file.

What happened?

The note is strict about what transfers. This CPU has no HBM, so the speedup here comes from cache traffic, not the HBM traffic the paper measures. The direction of the argument transfers, and the note quotes no GPU speedup from these numbers. The memory result is the one to keep: flash never materializes the square.

Decision: keep SDPA for every model from Day 2 onward. The Day 1 path stays as the reference the fused one is checked against. On the Day 4 T4, SDPA uses the memory-efficient kernel instead, because PyTorch's FlashAttention kernel needs a newer GPU.

Skipped

What we did not build, and why