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 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.
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 . For long sequences those reads and writes, not the multiplications, set the time. The paper counts 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.
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 , the sum , and the unnormalized output . A new block of scores arrives with maximum . Then:
The factor rescales everything computed under the old maximum to the new one. After the last block, is exactly the softmax-weighted average. No step needed the whole row.
What happened?
- The running maximum only moves up, and when it does, the running sum drops by the factor before the new block adds to it.
- The partial output after each block is already a valid attention output over the keys read so far. It converges to the full answer at the last block.
- Change the block size and the final output is the same to about . The block size is a performance choice, not a numerical one.
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:
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_sumIt 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.
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.
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.
What happened?
- The math backend's memory grows about fourfold per doubling, from 107 MiB at 1,024 tokens to 4,928 MiB at 8,192. That is the score matrix.
- The flash backend's memory grows with the inputs: 59 MiB at 8,192.
- At 8,192 tokens flash is 10.5 times faster in float32.
- bfloat16 is slower than float32 on this CPU for both backends. That is the opposite of a GPU, and the note flagged it for Day 4's mixed-precision work.
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.
What we did not build, and why
- A CUDA kernel. The tiled sketch teaches the blocking. Writing the real kernel teaches CUDA, which is a different project.
- The backward pass. FlashAttention recomputes the scores during backpropagation instead of storing them. The sketch is forward-only.
- GPU numbers for this benchmark. It reports process resident memory, which does not describe GPU allocation. Running it on the T4 would produce a number that looks like a measurement and is not one.