Skip to content

Fix FA4 b19 compatibility and drop the MLA-specific FA4 wrapper - #65

Merged
MasterJH5574 merged 2 commits into
mlc-ai:mainfrom
haok1402:0625-check-mlafa4-compile
Jun 25, 2026
Merged

MasterJH5574 merged 2 commits into
mlc-ai:mainfrom
haok1402:0625-check-mlafa4-compile

Conversation

@haok1402

@haok1402 haok1402 commented Jun 25, 2026 •

Copy link
Copy Markdown
Collaborator

Issues resolved

  1. 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 with ValueError: too many values to unpack. Fixed by discarding the extra values at the call site; floor bumped to flash-attn-4>=4.0.0b19.

  2. MLA-specific FA4 wrapper removed. DeepSeek-V2-Lite's non-CP MLA attention used an opaque mla_flash_attn_func custom 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 to flash_attn_func directly, so it lives in the torch.compile graph. The wrapper (_mla_fwd/_mla_bwd/mla_flash_attn_func) is deleted.

Correctness (B200, pp2/ep2, real routing) — no NaN

2026-06-25 13:46:45 | INFO | step 00000001/00000025 | step-time 57.087 sec | cross-entropy-loss 11.9269 | load-balance-loss 1.035905 | learning-rate 1.000000e-06 | gradient-norm 44.0787 | tokens-per-second 2,296 | peak-gpu-memory 78.11 GB
2026-06-25 13:46:47 | INFO | step 00000002/00000025 | step-time 2.546 sec | cross-entropy-loss 11.9132 | load-balance-loss 1.034085 | learning-rate 1.000000e-06 | gradient-norm 68.2542 | tokens-per-second 51,484 | peak-gpu-memory 78.12 GB
2026-06-25 13:46:50 | INFO | step 00000003/00000025 | step-time 2.488 sec | cross-entropy-loss 11.9035 | load-balance-loss 1.033514 | learning-rate 1.000000e-06 | gradient-norm 67.4410 | tokens-per-second 52,671 | peak-gpu-memory 78.12 GB
2026-06-25 13:46:53 | INFO | step 00000004/00000025 | step-time 2.491 sec | cross-entropy-loss 11.9059 | load-balance-loss 1.032568 | learning-rate 1.000000e-06 | gradient-norm 73.6298 | tokens-per-second 52,610 | peak-gpu-memory 78.12 GB
2026-06-25 13:46:56 | INFO | step 00000005/00000025 | step-time 2.493 sec | cross-entropy-loss 11.8912 | load-balance-loss 1.032724 | learning-rate 1.000000e-06 | gradient-norm 47.0430 | tokens-per-second 52,580 | peak-gpu-memory 78.12 GB

Performance (H100, throughput benchmark)

Moving the MLA concat out of the opaque custom op and into the torch.compile graph lets Inductor fuse it with the surrounding attention ops.

Before (concat boxed in the custom op):

2026-06-23 01:30:40 | INFO | step 00000001/00000025 | step-time 62.105 sec | cross-entropy-loss 11.9881 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 35.6103 | tokens-per-second 33,768 | peak-gpu-memory 39.73 GB
2026-06-23 01:31:04 | INFO | step 00000002/00000025 | step-time 23.277 sec | cross-entropy-loss 11.9855 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 39.1347 | tokens-per-second 90,094 | peak-gpu-memory 47.25 GB
2026-06-23 01:31:28 | INFO | step 00000003/00000025 | step-time 23.475 sec | cross-entropy-loss 11.9792 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 35.4346 | tokens-per-second 89,337 | peak-gpu-memory 47.26 GB
2026-06-23 01:31:51 | INFO | step 00000004/00000025 | step-time 22.993 sec | cross-entropy-loss 11.9767 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 34.9587 | tokens-per-second 91,207 | peak-gpu-memory 47.26 GB
2026-06-23 01:32:15 | INFO | step 00000005/00000025 | step-time 23.685 sec | cross-entropy-loss 11.9684 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 33.3818 | tokens-per-second 88,544 | peak-gpu-memory 47.26 GB

After (concat in the compiled graph):

2026-06-25 14:19:07 | INFO | step 00000001/00000025 | step-time 57.900 sec | cross-entropy-loss 11.9897 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 36.6775 | tokens-per-second 36,220 | peak-gpu-memory 39.72 GB
2026-06-25 14:19:30 | INFO | step 00000002/00000025 | step-time 21.952 sec | cross-entropy-loss 11.9874 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 29.3726 | tokens-per-second 95,533 | peak-gpu-memory 47.25 GB
2026-06-25 14:19:52 | INFO | step 00000003/00000025 | step-time 22.029 sec | cross-entropy-loss 11.9779 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 28.8567 | tokens-per-second 95,201 | peak-gpu-memory 47.25 GB
2026-06-25 14:20:15 | INFO | step 00000004/00000025 | step-time 21.900 sec | cross-entropy-loss 11.9726 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 28.6510 | tokens-per-second 95,760 | peak-gpu-memory 47.26 GB
2026-06-25 14:20:37 | INFO | step 00000005/00000025 | step-time 21.900 sec | cross-entropy-loss 11.9699 | load-balance-loss 1.000000 | learning-rate 1.000000e-06 | gradient-norm 37.7204 | tokens-per-second 95,759 | peak-gpu-memory 47.26 GB

Median throughput over steps 3–5 (warmup/compile excluded):

tokens/sec
Before 89,337
After 95,759
Speedup +7.2%

…n ; Drop MLA-specific FA4 wrapper and concat q/k outside the custom op

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread pithtrain/operators/flash_attn_v4.py
@github-actions

Copy link
Copy Markdown

Code review

1 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:
'Supports both symmetric (GQA/MHA) and asymmetric (MLA) head dimensions under BSHD layout.'
to:
'Supports symmetric (GQA/MHA) head dimensions under BSHD layout.'

See:

Wraps FA4's internal _flash_attn_fwd/_flash_attn_bwd with torch.library.custom_op
so that torch.compile can trace through them. Supports both symmetric (GQA/MHA)
and asymmetric (MLA) head dimensions under BSHD layout.
"""


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.

@haok1402

Copy link
Copy Markdown
Collaborator Author

Code review

1 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: 'Supports both symmetric (GQA/MHA) and asymmetric (MLA) head dimensions under BSHD layout.' to: 'Supports symmetric (GQA/MHA) head dimensions under BSHD layout.'

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.

@MasterJH5574 MasterJH5574 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM!

@MasterJH5574
MasterJH5574 merged commit 27c1264 into mlc-ai:main Jun 25, 2026
5 checks passed
@haok1402
haok1402 deleted the 0625-check-mlafa4-compile branch June 29, 2026 18:09

This branch was previously deployed

1 inactive deployment
review — 5f453128 Deployed Jun 25, 2026 by haok1402 via claude-review #11
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants