deepseek-ai / deepseek-ai/FlashMLA

10 rounds of hand-written fused attention that never beat a 3-kernel cuBLAS stack — trajectory, where each round died, and a correctness trap

Open
#210 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
C++
Stars
12.9k
Forks
1.2k
Avg merge
4h 20m
Merged PRs (30d)
2

Description

**Title:** 10 rounds of hand-written fused attention that never beat a 3-kernel cuBLAS stack — trajectory, where each round died, and a correctness trap

Not a bug report. We spent a campaign trying to beat a cuBLAS-based attention stack with a single fused kernel, got to 1.149× *slower* and stopped. The project is shelved, so I'd rather the round-by-round data be useful to people who do this for a living than sit in a dead repo. Close this if it's off-topic.

## The trajectory (RTX 5070, sm_120, 48 SMs, uncontended)

Baseline throughout = `cublasGemmStridedBatchedEx` QK^T → write S to HBM → standalone softmax kernel → `cublasGemmStridedBatchedEx` ·V. Ours = one fused kernel, online softmax, S never materialized.

| round | design | best ratio vs cuBLAS stack |
|---|---|---|
| R3 | scalar fused | 9–15× slower |
| R4 | + register-resident O attempt | 3.4–5.0× slower |
| R5 | — | 3.47× slower |
| R7 | (regression) | 5.3–6.9× slower |
| R8 | true FA-2 algorithm, online softmax `m_i`/`l_i` carried **across** K/V tiles, BK=16, grid.x = N/16 | 1.60–1.99× slower |
| R9 | single-warp | 1.662× @N=4096 |
| **R10** | **multiwarp + `cp.async`, reg-S + smem-scratch-reduce softmax** | **1.149× @N=4096 — campaign best, still slower** |

The structural claim held the whole time — 1 launch, zero S↔HBM traffic, ratio 1.25+N/d less memory. **The wall never crossed 1.0×.** `F-FUSION-ATTN-FLASHALGO-ABOVE-CUBLAS`, `F-FUSION-ATTN-MULTIWARP`

Two specific things that cost us rounds:

- **Single-CTA tensor-core fusion is a trap.** Keeping the S tile in shared memory in one CTA gives you a beautiful structural story and **9.4–15.5× slower** wall — it uses 1 of 48 SMs. Occupancy beat the memory-traffic win by an order of magnitude. `F-FUSION-ATTN-WMMA-WALL`
- **Warp-shuffle softmax was falsified** as a design (R9 pre-registration). What actually worked was reg-S + smem-scratch reduction, which we had assumed would be slower. `F-FUSION-ATTN-FLASHALGO-REGS`

## The correctness trap — the part I'd most want someone else to avoid

Separately, we built a fused single-token decode kernel. It emitted valid PTX, passed `ptxas`, launched, ran, and produced finite plausible-looking floats. Against an f64 CPU reference it had **max relative error 1.61** against a 1e-2 tolerance. Reading the emitted PTX back:

- the online softmax over QK^T was **omitted entirely** — the attention-weighted PV used `inv_sqrt_hd` as a "deterministic-finite stand-in" weight
- the output GEMV read `attn_out` from a shared-memory region never written with real attention output
- RMSNorm accumulated via `atom.shared.add` into a shared slot **not zeroed at kernel entry**

And the 17-kernel eager path we were benchmarking *against* was also numerically hollow (rel ~1.9): placeholder unit-scale RMSNorm, raw `ex2.approx` softmax with no max-subtract, no normalization.

**Both arms of the benchmark were wrong, and both compiled cleanly.** It was also 40–42× slower than the thing it was supposed to beat. `F-FUSION-AUTOREGRESSIVE-DECODE-TIMED`

Structural emits pass every gate a compiler engineer naturally builds — parse, lower, emit, ptxas, launch, no NaN. If your kernel CI doesn't gate on numerics against a reference implementation from the very first kernel, a hollow kernel will sit in your benchmark table looking fast.

## Questions, if anyone has a spare minute

1. **Is 1.0× against a cuBLAS batched stack the right bar at all**, or is the real win only at decode/batch-1 where the launch-bound regime dominates and the batched GEMM path is not the relevant baseline? We chose the batched-GEMM baseline early and never revisited whether it was the honest opponent.
2. **Did MLA's low-rank KV change where the fusion boundary sits** for you? Our fusion boundaries kept landing at ops needing a full row (softmax normalization), and I'm curious whether the latent projection moves that or just relocates it.

Happy to point at raw verdict files. Full post-mortem (MIT, includes three headline numbers we got wrong and retracted): https://github.com/dancinlab/hexa-lang/issues/5035

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.