Case Study: Steel Attention
steel_attention.h (~476 lines) is the complete
flash attention forward built from the
steel vocabulary: the most readable production FA on the
platform, and the payoff for knowing online softmax
before you arrive.
Specialization first: the function-constant set that compiles masking and raggedness out of pipelines that don't need them.
constant bool align_Q [[function_constant(200)]];
constant bool align_K [[function_constant(201)]];
constant bool has_mask [[function_constant(300)]];
constant bool do_causal [[function_constant(301)]];
constant bool has_sinks [[function_constant(302)]];
The register layout tells you the loop's shape: Q resident, K/V streaming, and
Stile (the score tile that naive attention writes to DRAM) as pure
register state.
MMATile<AccumType, TQ, 1, MMAFrag_acc_t> Qtile;
MMATile<AccumType, 1, TK, MMAFrag_acc_t> Ktile;
MMATile<AccumType, TQ, TK, MMAFrag_acc_t> Stile;
MMATile<AccumType, 1, 1, MMAFrag_acc_t> Vtile;
MMATile<AccumType, TQ, TD, MMAFrag_acc_t> Otile;
Stile is born from Q·Kᵀ, masked in place, softmaxed in place, multiplied
against V, and dies without ever touching memory. This in-place transformation is
why attn/ forks its own mma.h with a different fragment
layout.
Then the derivation, line for line, inside the
KV loop (ExpSubOp::apply(x,y) = fast::exp2(x - y); base-2 via scale *= M_LOG2E_F at line 166↗):
// Row max
Stile.template row_reduce<MaxOp>(new_max);
// exp(Si - rowmax(Si))
Stile.template row_bin_op<ExpSubOp>(new_max);
// Factor exp(rowmax(Si) - rowmax(Si-1))
for (short i = 0; i < kRowsPT; ++i)
factor[i] = fast::exp2(max_score[i] - new_max[i]);
...
// Update norm
sum_score[i] = sum_score[i] * factor[i] + sum_score_tmp[i];
// Update O
Otile.template row_bin_op<MulOp>(factor);
— steel_attention.h:391-420↗, abridged
factor is the correction c; the final DivOp normalize is at
line 460↗.
If the online-softmax page landed, this file
holds no surprises, which is the point of reading it second.
Details that reward attention: causal handling is a loop-bound computation, not
a mask. kb_lim at
lines 239-247↗
means tiles above the diagonal are never visited
(question 1: delete work). K loads
transposed via a differently-parameterized
BlockLoader. GQA is a stride trick (kv_head_idx = tid.y / gqa_factor). What this kernel serves:
mx.fast.scaled_dot_product_attention's prefill path, for
the head dims its dispatch gate accepts. Decode goes
elsewhere; backward
doesn't exist on Metal.