FlashAttention forward in C
Overview
Section titled “Overview”| Module | L9.3 · side · C · Pass 6 · 6 to 8 h |
| You build | c/src/kernels/flash_attn.c: tl_flash_attn_fwd_f32, attention for a whole prefill (or a chunk of one) that never stores the score matrix: GQA head sharing, causal masking at absolute positions (q_offset), a sliding window, learned sinks, and the log-sum-exp of every row, with scratch from your arena and threads from your pool |
| Contract | course/contracts/c/include/tinyllm/attention.h · rules: c/ABI.md (rule 10, invariance) |
| Tests | course/tests/L9.3/: test_flash_attn.c (C, under ASan and UBSan, and ThreadSanitizer for the pool case, against a naive oracle built on your L9.2 softmax) and shared file fixtures (what they check: section 4) |
| Needs | rt.02 the loader · rt.02 the arena · rt.03 the pool · M09.6 tl_expf · L9.2 tl_softmax_f32 (the C oracle) · L7.7 windowed_attention · L5.1 sdpa_forward (or --ref-deps). Reading: L5.2 (mask flags), L7.5 (GQA) |
| Used by | L9.4 checks paged decode against this kernel, bit for bit · L10.3 relies on the chunk invariance proved here |
| Milestone | MS-L9 |
| Optional depth | Dao et al., “FlashAttention” (2022) and “FlashAttention-2” (2023); Milakov and Gimelshein, “Online normalizer calculation for softmax” (2018); Rabe and Staats, “Self-attention Does Not Need Memory” (2021) |
Key Takeaways
Section titled “Key Takeaways”- Attention can be computed one tile of keys at a time with the online softmax of
L9.2, extended to the output: each row keeps a running maximum , denominator , and output , and rescales and by when the maximum grows (gqa_causal_window_match_naive). - Memory stops growing with the context: scratch is rows of state and one tile of scores per worker, so the arena’s high-water mark is the same for 16 keys and 1000 (
arena_rewound_and_scratch_independent_of_tk). - Visibility is per row and per absolute position: query sits at and sees key when (causal) and (window). Rows of one query tile may see different keys (
no_visible_key_gives_zeros_and_minus_inf). - Key tiles are aligned to absolute positions (), so a row’s arithmetic is the same whether the prompt arrives in one call or in chunks, in any batch, with any query tile size: bitwise (
chunk_invariant_with_q_offset,query_tile_and_batch_invariant). - A sink is one extra score with no value: it joins the denominator and the lse, and the weights then sum to (
sinks_join_the_denominator).
How to work this chapter
Section titled “How to work this chapter”ol start L9.3 # stubs c/src/kernels/flash_attn.c into your repool tests L9.3 # read the test catalog firstol check L9.3 # exit code is the verdictol check L9.3 --ref-deps # only if a runtime piece, L9.2, L7.7, or L5.1 is not passing yetol parity flash.fwd # against the float64 goldenol diff L9.3 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Your numpy attention (L5.1, L7.5, L7.7) forms the score matrix for every head: floats, then a softmax over each row, then . For a 2048-token prompt that is 4 million floats (16 MB) per head per layer, written once and read twice, and most of the prefill time goes to moving it rather than computing it. This standalone C exercise makes memory use independent of prompt length. L10.3 additionally splits long prompts into chunks so prefill does not stall decode; chunk invariance ensures every query row’s output is stable across split choices.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| query row , key and value at position | float[D] | |
| batch, query heads, KV heads () | int64_t | |
| query rows in this call, keys in the cache | int64_t | |
| the absolute position of query row | int64_t | |
scale, usually | float | |
| a score | float | |
| the keys row may see | set | |
| the window (: none) | int64_t | |
| the sink logit of head | float | |
| query and key tile sizes | int64_t | |
| running maximum, denominator, and output of row | float, float, float[D] | |
float |
2.1 What is computed
Section titled “2.1 What is computed”For every sequence , head , and query row :
The term is present only with sinks. GQA: head reads KV head , so the query heads of a group share one KV head (L7.5). Positions are absolute: row of a chunk that starts at token 24 is at , and key is at position , so the same flags serve a prefill (, ), a chunk of a prefill, and a decode step (, ). A row with no visible key and no sink has no distribution: the contract says and .
2.2 The tiled online softmax
Section titled “2.2 The tiled online softmax”L9.2 showed that a running maximum and a running sum can absorb a row’s entries one block at a time: when a block raises the maximum to , multiply the old sum by . The same factor fixes the running output , because each of its terms was weighted against the old maximum. For one key tile with visible scores :
After the last tile, and . A sink is the cheapest possible tile: start each row at , , . Without a sink, start at , ; the first visible tile then has , which multiplies a zero accumulator. A tile in which the row sees nothing must be skipped for that row, not processed: with and no finite score, is .
2.3 Tiles, scratch, and why memory stops growing
Section titled “2.3 Tiles, scratch, and why memory stops growing”The kernel loops over work items (sequence, head, query tile of rows). For each item it keeps for rows, floats, and one tile of scores, and walks the key tiles that any of its rows can see. Each key tile is read from memory once per query tile and used by all rows while it is in cache, which is FlashAttention’s whole point on a GPU (where the tiles live in shared memory) and helps on a CPU too. The scratch per worker is floats, from the caller’s rt.02 arena (one slice per worker, taken before the threads start, because the arena is not thread-safe). None of it depends on . The kernel takes an arena mark on entry and rewinds to it before returning, so the caller’s arena is exactly as it was; with scratch == NULL it uses a private arena and destroys it.
2.4 Invariance: key tiles at absolute positions
Section titled “2.4 Invariance: key tiles at absolute positions”Key tile always covers positions , whatever the query tile, the chunk, or the batch. Within the walk, each row:
- computes its own (its own
loandhi), - skips every tile that holds none of its keys,
- in a tile that holds some, adds its visible keys in increasing .
So the sequence of floating-point operations that produces row depends only on , the keys and values, , the flags, and . Not on which other rows share its query tile (), not on how the prompt was chunked ( moves and the row together), not on the batch, and not on the worker. That is the chunk invariance L10.3 needs: query rows computed in one call or in chunks with q_offset are bitwise equal. Two tempting shortcuts break it: starting the tiles at a query tile’s first visible key (the tile boundaries then move with the chunk), and deciding visibility once per query tile instead of per row.
2.5 Threads
Section titled “2.5 Threads”Work items are independent: tl_parallel_for runs them on rt.03 workers, each using the scratch slice of its worker index. Every row is computed by one worker with the arithmetic above, so 4 threads give the serial bits.
3. Worked example by hand
Section titled “3. Worked example by hand”One query , three keys and values, no mask, , :
| 0 | 1 | ||
| 1 | 0 | ||
| 2 | 1 |
Tile 0 (keys 0, 1). Start , , . Tile maximum 1, so and . Weights and :
Tile 1 (key 2). Tile maximum 1, , (no rescale). Weight :
Normalize: exactly (because and ), and .
This is hand_example and test_hand_example. With a causal mask and the row would see only keys 0 and 1 (tile 0), giving .
4. The interface
Section titled “4. The interface”/* tinyllm/attention.h: q, o [B, H, Tq, D]; k, v [B, Hkv, Tk, D]; lse [B, H, Tq] or NULL */tl_status tl_flash_attn_fwd_f32(const float *q, const float *k, const float *v, float *o, float *lse, int64_t B, int64_t H, int64_t Hkv, int64_t Tq, int64_t Tk, int64_t D, float scale, int64_t q_offset, int causal, int64_t window, const float *sink_logits /* [H] or NULL */, int64_t Br, int64_t Bc, tl_arena *scratch /* or NULL */, tl_pool *tp /* or NULL */);/* Br, Bc = 0 pick 64. TL_EINVAL: negative sizes or q_offset, NULL tensors. TL_ESHAPE: H % Hkv != 0 or Hkv > H. TL_ENOMEM: the arena cannot grow. */What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
hand_example | unit, smoke | section 3: , lse | the definition and the lse |
gqa_causal_window_match_naive | differential | GQA 6:2 and MQA, causal prefill, a chunk at q_offset 24, window 6, window without causal, plain; odd tiles and the default 64 | every flag of a Llama-family model |
sinks_join_the_denominator | unit | sinks against the oracle; all-ones values give and lse | gpt-oss style learned sinks |
no_visible_key_gives_zeros_and_minus_inf | boundary | a query whose window excludes every key: zeros and ; with a sink, zeros and lse = the sink | no NaN in the residual stream |
chunk_invariant_with_q_offset | property | 37 rows in one call vs chunks of 1, 3, 5, 16 with q_offset, bitwise | chunked prefill (L10.3) |
query_tile_and_batch_invariant | property | bitwise; a sequence alone vs in a batch of 3 | batching (L10.2) |
arena_rewound_and_scratch_independent_of_tk | property | the arena is back at its mark; high water equal for and 1000 | memory flat in the context length |
pool_result_equals_serial_bitwise | property | 4 threads vs serial (also under TSan) | threads never change bits |
shape_and_argument_errors | boundary | TL_ESHAPE for 6 heads on 4 KV heads and ; TL_EINVAL cases leave alone; gives zeros | errors before any write |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| not rescaling the output accumulator when the maximum grows | earlier tiles are overweighted by | gqa_causal_window_match_naive (mutant s01) |
| not rescaling the denominator | outputs do not average the values | gqa_causal_window_match_naive (mutant s02) |
| mapping query heads to KV heads round-robin () | GQA models attend with the wrong keys; MHA still passes | gqa_causal_window_match_naive (mutant s03) |
| causal as | a token cannot see itself | gqa_causal_window_match_naive (mutant s04) |
| a window of keys | drifts from the model’s training | gqa_causal_window_match_naive (mutant s05) |
ignoring q_offset | a chunk at position 24 behaves like the start of a prompt | chunk_invariant_with_q_offset (mutant s06) |
| leaving the sink out of the denominator | weights sum to 1; the sink does nothing | sinks_join_the_denominator (mutant s07) |
| lse without the maximum ( instead of ) | lse off by ; any later merge of partial results breaks | hand_example (mutant s08) |
| dividing by | a row that sees nothing becomes NaN | no_visible_key_gives_zeros_and_minus_inf (mutant s09) |
| key tiles starting at a query tile’s first visible key | the bits depend on the chunking | chunk_invariant_with_q_offset (mutant s10) |
| not rewinding the caller’s arena | the engine’s step arena grows every call | arena_rewound_and_scratch_independent_of_tk (mutant s11) |
| a score buffer sized by | scratch grows with the context: the memory FlashAttention exists to save | arena_rewound_and_scratch_independent_of_tk (mutant s12) |
| a pooled partition that drops a work item | wrong only with threads | pool_result_equals_serial_bitwise (mutant s13) |
| deciding visibility once per query tile | rows in one tile see the wrong keys; results depend on | query_tile_and_batch_invariant (mutant s14) |
| accepting | head groups read past the KV heads | shape_and_argument_errors (mutant s15) |
pointer arithmetic on a NULL k when | UBSan: “applying zero offset to null pointer” | shape_and_argument_errors |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | rt.02 | the error slot and allocator support |
| Back | rt.02 | the arena that holds the tiles, with a mark taken and rewound around each call |
| Back | rt.03 | tl_parallel_for over (sequence, head, query tile) |
| Back | M09.6 | tl_expf for every weight and rescale factor |
| Back | L9.2 | the online softmax update this kernel applies to tiles; tl_softmax_f32 is the C tests’ oracle |
| Back | L7.7 | windowed_attention, the specification of every flag |
| Back | L5.1 | sdpa_forward, the specification without a causal mask |
| Forward | L9.4 | paged attention for decode performs the same tile update on blocks; its tests hold it to this kernel bit for bit |
| Forward | the standalone Rust engine | (Pass 7) the Rust forward’s prefill; L10.3’s chunked prefill relies on the chunk invariance |
If you use the optional C paged-attention exercise, its L9.3 prerequisite supplies this standalone C kernel.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
| CPU tiles in an arena | FlashAttention-2 (CUDA) | tiles in shared memory, warps split over queries, the backward pass by recomputation | Dao-AILab/flash-attention, csrc/flash_attn/src/ |
| one kernel for all flags | FlashInfer | variants generated per mask, layout, and positional encoding; paged KV inputs | flashinfer/include/flashinfer/attention/ |
| float32 scores | FlashAttention-3 | fp8 inputs with incoherent processing, asynchronous tiles on Hopper | Shah et al. (2024) |
| fixed tile order for invariance | split-KV (FlashDecoding) | splits long contexts across thread blocks and merges partial (m, l, a) with the same rescale, trading invariance for parallelism | Dao et al., “Flash-Decoding for long-context inference” (2023) |
| float32 KV only | llama.cpp ggml_flash_attn_ext | CPU and GPU flash attention over quantized KV | ggml/src/ggml-cpu/ops.cpp |