Repository navigation
Conversation
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.
Summary
This PR adds optional per-request
topk_lengthsupport to the SM90DeepSeek-V3.2 FP8 sparse decode kernel (
d_qk=576). SM90 V3.2 can now usedynamic top-k instead of always processing the full allocated top-k capacity.
When
topk_lengthis provided, the kernel:max(ceil_div(topk_length, TOPK_BLOCK_SIZE), 1)blocks;When
topk_length == nullptr, the existing fixed top-k behavior is preserved.There is no public API change, and MODEL1 and SM100 behavior are unchanged.
Why
template<bool DYNAMIC_TOPK, ...>I also tested a single-device-kernel implementation that checks
params.topk_length != nullptrat runtime. The condition is grid-uniform, butthe compiler cannot remove the length load, tail predicates, extra control
flow, or their value lifetimes from the fixed path.
template<bool DYNAMIC_TOPK>moves the choice to host launch dispatch:false: the dynamic-length logic is compiled out, preserving the fixed path;true: per-row block counts and partial-tail safety checks are enabled.The comparison implementation is available at
7b7c831ontest/sm90-v32-dynamic-topk-runtime-device. It regresses the representativeB64-B128 fixed cases by 8.4%-9.4%, while the specialized implementation stays
within 0.5% of the baseline. Both builds report 168 registers, 12 barriers,
and no spills, so the generated hot-path instructions—not an occupancy
change—explain the observed difference.
Correctness
Focused reference checks passed for:
0,1,17,63,64,65,127,128For the tail-safety check, every suffix entry after
topk_lengthwas replacedwith
INT_MAX. H64 and H128 matched the corresponding safe-tail results, withno illegal memory access.
Kernel performance
Environment: NVIDIA H100 80GB HBM3 (SM90), driver 580.126.09, nvcc 12.8.93,
and PyTorch 2.13.0+cu129. Each number is the median of seven group means, with
10 warmups and 20 timed calls per group using CUDA events. Scheduler metadata
is reused and excluded from timing.
These are standalone FlashMLA sparse-decode measurements; they exclude input
generation, index compaction, framework integration, collectives, and model
execution. No end-to-end speedup is claimed.
Compared revisions:
15f13e57b7c831cfc99c2Fixed top-k regression check
Shape:
S_q=2,H=128,d_qk=576,s_kv=32768, page block 64,variable KV lengths,
topk=2048. Positive deltas are regressions.For B64-B128, the runtime-check version regresses by 8.97% geomean. The final
specialized version is -0.25% versus baseline, showing no fixed-path
regression within measurement variation.
Dynamic top-k
Shape:
[B=256, S_q=1, H=128, d_qk=576],s_kv=4096, page block 64,allocated top-k 2048, effective per-row top-k 1024. The fixed input uses an
invalid suffix; the dynamic input passes
topk_length=1024and poisons theignored suffix with
INT_MAX.The final dynamic path saves 130.89 us per call versus its fixed path. It is
also 6.90% faster than the shared runtime-check dynamic path, while
specialization avoids the latter's 8.51% fixed-path regression in this shape.
Reproduction
Build each worktree with:
Save the folded harness below as
bench_sm90_v32_dynamic_topk.py, thenrun it in a fresh process for each build:
Benchmark harness (click to expand)
Scope
This change only enables dynamic top-k for the existing SM90 V3.2 sparse
decode specialization. It does not compact selected indices; callers that
provide
topk_lengthmust place valid indices in a prefix. Framework-sidelayout conversion and integration should be benchmarked separately.