Fix FA4 b19 compatibility and drop the MLA-specific FA4 wrapper - #65
Conversation
…n ; Drop MLA-specific FA4 wrapper and concat q/k outside the custom op
There was a problem hiding this comment.
Code Review
This pull request removes the custom MLA flash attention implementation (mla_flash_attn_func and its custom ops) and replaces it with standard flash_attn_func by manually concatenating the query and key projections. It also updates the flash-attn-4 dependency to >=4.0.0b19. Feedback highlights a critical issue where other call sites of _flash_attn_fwd in pithtrain/operators/ring_attention.py still unpack only two values, which will cause a runtime crash with the upgraded dependency.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
Code review1 issue found pithtrain/operators/flash_attn_v4.py, line 6 — Stale docstring references deleted MLA code The module docstring says it supports asymmetric (MLA) head dimensions, but this PR deleted all MLA-specific code (mla_flash_attn_func, _mla_fwd, _mla_bwd). Per AGENTS.md: avoid adding removed comments for removed code. The fix is to remove the MLA reference from line 6, changing: See: pith-train/pithtrain/operators/flash_attn_v4.py Lines 4 to 8 in 5f45312 Everything else looks good: the unpacking fixes are correct at all call sites, the MLA wrapper deletion is complete with no dangling references, and the inline Q/K concat correctly replicates the deleted op via autograd. |
Wrong... We removed the MLA wrapper but at the same time FA4 already supports MLA so that docstring above is accurate and to-the-point. |
Issues resolved
FA4 b19 4-tuple return. flash-attention b19 ([Cute,Bwd,Sm100] add sparse MLA (Deepseek v4) backward kernels Dao-AILab/flash-attention#2621, sparse MLA) widened
_flash_attn_fwd's return from(out, lse)to(out, lse, p, row_max). Every call site in the repo unpacked two values and crashed withValueError: too many values to unpack. Fixed by discarding the extra values at the call site; floor bumped toflash-attn-4>=4.0.0b19.MLA-specific FA4 wrapper removed. DeepSeek-V2-Lite's non-CP MLA attention used an opaque
mla_flash_attn_funccustom op that boxed the q/k concat + FA4 together, to dodge an Inductor SM100 codegen NaN on FA4's asymmetric-dim backward (torch 2.10 + FA4-b7). That NaN no longer reproduces on the current torch 2.12 + FA4-b19 stack. The concat is now done in the model and fed toflash_attn_funcdirectly, so it lives in thetorch.compilegraph. The wrapper (_mla_fwd/_mla_bwd/mla_flash_attn_func) is deleted.Correctness (B200, pp2/ep2, real routing) — no NaN
Performance (H100, throughput benchmark)
Moving the MLA concat out of the opaque custom op and into the
torch.compilegraph lets Inductor fuse it with the surrounding attention ops.Before (concat boxed in the custom op):
After (concat in the compiled graph):
Median throughput over steps 3–5 (warmup/compile excluded):