Conversation
FlashAttentionForwardSm120 fails at dispatch on real SM120 hardware (validated on SM121a / DGX Spark GB10) with three distinct symptoms that all stem from one root cause. FlashAttentionForwardBase.__init__ assigns self.arch = BaseDSL._get_dsl().get_arch_enum(), which on SM120 returns sm_121a and overwrites the class level arch = 80 declaration on FlashAttentionForwardSm120. Add an __init__ override that runs after super().__init__() and restores self.arch = Arch.sm_80, sets self.is_split_kv = False, and sets self.pack_gqa = False. All changes are contained to FlashAttentionForwardSm120. Validated on SM121a: bf16 and fp16, causal and non causal, dense and varlen, six shape configs: 48 / 48 pass with max abs diff <= 0.0156 against a PyTorch fp32 reference. GQA configs pass through the non packed path. Signed-off-by: Blake Ledden <blake@secondnaturecomputing.com> Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This was referenced Apr 22, 2026
3 tasks
|
Bump, master still breaks on sm_120. @tridao can you review this please? |
|
Any plans/updates? |
Contributor
Author
|
Closing this PR. The three failure modes it addressed have since been fixed on main by other means: the SM120 shim now forces the SM80 arch path after |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
FlashAttentionForwardSm120fails at dispatch on real SM120 hardware (validated on SM121a / DGX Spark GB10) with three distinct symptoms that all stem from one root cause.FlashAttentionForwardBase.__init__assignsself.arch = BaseDSL._get_dsl().get_arch_enum(), which on SM120 returnssm_121aand overwrites the class levelarch = 80declaration onFlashAttentionForwardSm120. Three symptoms follow:AttributeError: 'FlashAttentionForwardSm120' object has no attribute 'is_split_kv'on every call. The base__init__does not acceptis_split_kv(only Sm90 and Sm100 do), but the sharedSm80.__call__readsself.is_split_kvdirectly.self.use_tma_O = self.arch >= Arch.sm_90resolves True onsm_121a, so the TMA O branch runs and fails withAttributeError: 'NoneType' object has no attribute '_trait'atcpasync/helpers.py:209(TMA atom is None because SM120 has no TMA O support).ragged = self.use_tma_O and has_cu_seqlensis True for the same reason, hitting an incompatible layout path and failing atflash_fwd.py:398withexpects coord and shape of view are weakly congruent, but got '!cute.layout<"(?,?):(?{i64 div=8},1)">', '!cute.coord<"(_,_,?)">'.Fix: add an
__init__override onFlashAttentionForwardSm120that runs aftersuper().__init__()and restoresself.arch = Arch.sm_80, setsself.is_split_kv = False, and setsself.pack_gqa = False.The
pack_gqaoverride addresses a separate preexisting issue.Sm80.__call__'s pack GQA path does not callpack_gqa_layoutbefore handing tensors toPackGQAmethods, whileSm90.__call__(L273) andSm100.__call__(L504) do. Porting those three lines intoSm80.__call__verbatim is not sufficient, because the Sm80 mainloop's tile sizing assumestile_mdivides the seqlen dimension cleanly, which fails for the packed(qhead_per_kvhead, seqlen_q)layout whenqhead_per_kvheaddoes not dividetile_m(for example,cute.local_tileraises on an unaligned division). A proper Sm80 pack GQA fix needs a follow up that also adjusts the tile scheduler. Forcingpack_gqa = Falseon the Sm120 subclass sidesteps both issues without touching the Sm80 class, keeping blast radius off SM8.All changes are contained to
FlashAttentionForwardSm120. Nothing outside that class is touched, so SM8 / SM9 / SM10 / SM11 dispatch paths are unaffected.Scope
This is a correctness fix that unblocks
FlashAttentionForwardSm120end to end. It is not a performance improvement. FA4's perf advantage over FA2 comes from tcgen05 and TMEM on SM100. SM120 has neither, soFlashAttentionForwardSm120compiles down to the same SM80 eramma.syncpath as FA2 with additional dispatch complexity on top. On the realistic Qwen 3 prefill shapes I measured, patched FA4 on SM120 runs within roughly 15 percent of FA2 at longer sequences and is slower at short sequences. Users on SM120 are better served by FA2 today for pure attention throughput. The motivations to land this fix anyway are:score_mod, block sparse masking (Add SM80/SM120 block-sparse forward attention support #2389), and split KV FlashDecoding (Add SM120 split-KV (FlashDecoding) with FP32 partial outputs #2336) become available on SM120 once the crash is resolved.Relationship to in flight SM120 feature PRs
The
self.arch = Arch.sm_80portion of this fix also appears in three open SM120 feature PRs (#2336 split KV, #2348 paged KV, #2389 block sparse), each of which adds its own__init__carrying that one line as part of its feature scope. I am filing this as a standalone bug fix because the arch issue alone blocks every SM120 dispatch path regardless of which feature is enabled, and a small focused review surface is easier to land quickly. If this PR merges first, those three feature PRs rebase to a one line diff. If one of them merges first, this PR rebases to drop theself.archline. I own all four branches and am happy to handle the rebase either way.Validation
SM121a (DGX Spark GB10), bfloat16 and float16, causal and non causal, dense and varlen, six shape configs: 48 / 48 pass, max abs diff <= 0.0156 against a PyTorch fp32 reference. GQA configs pass through the non packed path (Qwen style
H_q / H_kv = 16 / 2atD = 128, MQA4 / 1atD = 128, MHAD = 96).Known limitations
_is_fa4_supported()gate already excludes SM120, so end users are not impacted by default.interface.py:753. Add SM120 split-KV (FlashDecoding) with FP32 partial outputs #2336 adds full SM120 split KV support as a feature.head_size > 128not supported on SM120 for this kernel variant due to the 99 KB SMEM budget. Matches vLLM's existingfa_utils.pygate that routeshead_size > 128to FA2 on Blackwell.AI assistance
AI assistance was used to author this patch and its test plan. I validated each change on SM121a hardware and reviewed every line before pushing.
Contributed by Second Nature Computing (https://joinsecondnature.com)