Skip to content

[AMD][gfx1250] Retune sparse paged-prefill launch config for DSV4-Flash - #41382

Open
ankith117 wants to merge 4 commits into
sgl-project:mainfrom
ankith117:amd-gfx1250-sparse-prefill-launch-config
Open

ankith117 wants to merge 4 commits into
sgl-project:mainfrom
ankith117:amd-gfx1250-sparse-prefill-launch-config

Conversation

@ankith117

@ankith117 ankith117 commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

The Triton path in _sparse_attn_v4_paged_prefill_triton currently launches with a conservative configuration: block_h = 16, the AMD MFMA minimum tile, and no num_stages argument. These are safe defaults that work across architectures, and the file has no @triton.autotune, so they apply uniformly wherever the Triton kernel is selected.

On gfx1250 there is meaningful headroom in that choice. The kernel spends most of its time on scattered tl.load gathers through kv_indices, which benefits from software pipelining, and with index_n_heads = 64 a larger head tile lets a single CTA cover all heads of a token. Profiling DeepSeek-V4-Flash at ISL 8192 puts this kernel at 9.9% of prefill GPU time and the highest per-call cost in the trace (473 us), so it is worth tuning for this architecture.

This PR selects a gfx1250-specific launch configuration and leaves every other target unchanged.

Modifications

Select BLOCK_H=64, BLOCK_K=16, num_warps=4, num_stages=2 on gfx1250. BLOCK_H=64 equals index_n_heads, so one CTA covers every head of a token and the grid stops splitting on the head dimension.

Gated: gfx950 takes the OPUS kernel, but gfx1250 and NVIDIA both fall through to this Triton path, and these constants were only measured on gfx1250. Every other target keeps the previous behaviour exactly.

Microbenchmark vs the shipped config (H=64, D=512, topk=512), two independent sweeps:

T shipped tuned speedup
128 279 us 98 us 2.84x
512 919 us 440 us 2.09x
2048 3856 us 1709 us 2.26x

The optimum moves with token count — T <= 256 prefers (32,32,8,2), T >= 512 prefers (64,16,4,2). Measured from the scheduler log, 90.1% of prefill tokens arrive in batches with T >= 384 (358 batches at the chunked-prefill cap plus a tail of small remainders), so the config targets large T.

Hardcoded rather than @triton.autotune: an autotune key= would need a token-count bucket, since T changes on nearly every chunked prefill. Reasonable as a follow-up — several attention kernels here already use a _BUCKET key — but a larger change than this fix warrants.

Adds test/registered/kernels/ops/attention/test_dsv4_paged_prefill_launch_config.py, the first test for this kernel: each launch config against a dense torch reference; all configs against each other on identical inputs, with T in {1,5,64,512} and H in {64,48} so BLOCK_H=64 exercises the head-mask tail; a token whose index lists are entirely -1; and that the wrapper dispatches the constants its arch branch claims, checked for both branches so non-gfx1250 CI validates the gate rather than skipping it. 13 tests, all passing on gfx1250.

Accuracy Tests

GSM8K, 1319 questions, 5-shot, temperature 0:

accuracy invalid
baseline 0.929 0.000
this PR 0.932 0.000

Unchanged, as expected by construction: the online softmax accumulation is tile-size invariant, so block sizes change scheduling and not numerics.

Speed Tests and Profiling

MI455X (gfx1250), DeepSeek-V4-Flash, TP=1 + EAGLE(2/1/3), ISL 8192 / OSL 1024, random sequences. Median of runs 2-5 with the warmup run dropped.

Throughput, +5.38% geomean

conc baseline tuned delta
8 410.9 433.8 +5.6%
16 658.2 696.5 +5.8%
32 934.7 983.6 +5.2%
64 1293.3 1356.8 +4.9%

Latency, TTFT -8.1% geomean, TPOT -5.4%

conc TTFT base TTFT tuned delta TPOT base TPOT tuned delta
8 276.0 ms 257.7 -6.6% 16.91 ms 15.97 -5.6%
16 367.2 ms 333.8 -9.1% 21.89 ms 20.69 -5.5%
32 494.1 ms 453.8 -8.2% 30.24 ms 28.59 -5.5%
64 746.1 ms 684.5 -8.3% 47.62 ms 45.20 -5.1%

Total token throughput at conc 64: 12,694 tok/s. Run-to-run spread is 0.6-2.3%, so the deltas are several times the noise. GPU SCLK was logged on every run (2159-2163 MHz) to confirm no clock throttling contaminated the measurement.

Measured on a tree that already includes #40996, so this is the incremental gain on top of those tuned fp8 configs.

Checklist


CI States

Latest PR Test (Base): ❌ Run #37713471603
Latest PR Test (Extra): ❌ Run #37713471329
Latest PR Test (AMD ROCm 10): ❌ Run #37713471591

_sparse_attn_v4_paged_prefill_triton hardcodes block_h = 16 (the MFMA minimum
tile, not the optimum) and never passes num_stages, so a kernel dominated by
scattered tl.load gathers runs with no software pipelining.

On gfx1250 use BLOCK_H=64, BLOCK_K=16, num_warps=4, num_stages=2. BLOCK_H=64
equals index_n_heads, so one CTA covers every head of a token.

Microbenchmark vs the shipped config (H=64, D=512, topk=512):
  T=128    279 us -> 98 us    2.84x
  T=512    919 us -> 440 us   2.09x
  T=2048  3856 us -> 1709 us  2.26x

The optimum moves with token count, and 90% of prefill tokens arrive in
batches at the chunked-prefill cap, so the config targets large T.

End-to-end on MI455X, DeepSeek-V4-Flash TP=1 + EAGLE, ISL 8192 / OSL 1024,
median of runs 2-5:
  conc 8   410.9 ->  433.8 tok/s  +5.6%   TTFT -6.6%
  conc 16  658.2 ->  696.5 tok/s  +5.8%   TTFT -9.1%
  conc 32  934.7 ->  983.6 tok/s  +5.2%   TTFT -8.2%
  conc 64 1293.3 -> 1356.8 tok/s  +4.9%   TTFT -8.3%

Throughput +5.38% geomean, TTFT -8.1%, TPOT -5.4%, spread 0.6-2.3%. GSM8K
0.932 with 0% invalid, unchanged: online softmax is tile-size invariant, so
block sizes change scheduling and not numerics.

Gated to gfx1250. gfx950 takes the OPUS kernel, but gfx1250 and NVIDIA both
reach this Triton path, and these constants were only measured on gfx1250;
every other target keeps the previous behaviour.

Adds the first test for this kernel: each config against a dense torch
reference, all configs against each other (H=48 exercises the head-mask tail),
the all-sentinel row case, and that the wrapper dispatches what its arch
branch claims, checked for both branches.

@HaiShaw HaiShaw left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

@sogalin @1am9trash please evaluate the tuned parameters, including 8k/1k beyond.

# gather-bound kernel needs. BLOCK_H=64 == index_n_heads, so one CTA
# covers all heads of a token. Measured 2.8x/2.1x/2.3x at T=128/512/2048
# (geomean ~2.2x). Tuned for large T: 90% of prefill tokens arrive in
# batches at the chunked-prefill cap. Gated because non-gfx95 NVIDIA

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

this is gfx use only triton kernel, I think you meant non-gfx1250 vs. non-gfx95 NVIDIA

@ankith117 ankith117 Sep 27, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Yes. Thank you for pointing it out. Fixed it.

@HaiShaw

HaiShaw commented Sep 27, 2026

Copy link
Copy Markdown
Collaborator

@ankith117 lint to fix.

@ankith117

Copy link
Copy Markdown
Contributor Author

@HaiShaw lint fixed. Thank you.

@ankith117

ankith117 commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor Author

MI455X (gfx1250), DSV4-Flash, TP=1 + EAGLE(2/1/3), ISL 1024 / OSL 1024, random sequences, median of runs 2-5 with the warmup run dropped same protocol as the 8k/1k table above.

conc Throughput baseline Throughput tuned delta TTFT base TTFT tuned delta
8 418.4 422.2 +0.9% 249.5 ms 243.1 -2.6%
16 686.9 692.9 +0.9% 302.2 ms 296.1 -2.0%
32 976.4 988.9 +1.3% 414.6 ms 401.5 -3.2%
64 1299.7 1333.4 +2.6% 633.6 ms 610.5 -3.6%

Throughput +1.41% geomean, TTFT -2.86%, TPOT -1.62%. Run-to-run spread 0.3-3.2%; SCLK 2155-2366 MHz throughout, so no throttling.

The gain is smaller than at 8k/1k (+5.38%) and that is the expected direction: at ISL=OSL the run is decode-dominated, so a prefill-only kernel change has proportionally less to act on. The useful result is that there is no regression at any concurrency BLOCK_H=64 was tuned for large T, and short prompts were the plausible place for it to cost something.

@ankith117

Copy link
Copy Markdown
Contributor Author

MI455X (gfx1250), DSV4-Flash, TP=1 + EAGLE(2/1/3), ISL 70000 / OSL 300, random sequences, 5 runs per point, same harness as the 8k/1k and 1k/1k tables above.

At this context length the random dataset shares long prefixes, so the radix cache serves most of each prompt after the first run (#cached-token climbs past 100k; run 1 is ~20x slower than run 2 at every concurrency). Since the prefill kernel only runs on tokens that are not cache hits, the two passes answer different questions and I am reporting both.

Cache-cold (run 1 — the pass that actually prefills), n=1 per point:

conc baseline tuned delta TTFT base TTFT tuned delta
8 18.8 21.2 +13.2% 12955 ms 11028 -14.9%
16 36.5 41.5 +13.6% 6684 ms 5651 -15.4%
32 40.1 45.7 +14.0% 9043 ms 7135 -21.1%
64 41.0 46.8 +14.2% 9900 ms 8412 -15.0%

Throughput +13.77% geomean, TTFT -16.66%.

Cache-warm (runs 2-5, the standard median-of-4 protocol):

conc baseline tuned delta TTFT base TTFT tuned delta
8 375.5 375.5 -0.0% 321.3 ms 309.9 -3.5%
16 491.8 499.5 +1.6% 442.0 ms 418.6 -5.3%
32 739.3 752.5 +1.8% 615.3 ms 582.8 -5.3%
64 938.6 962.9 +2.6% 961.5 ms 904.1 -6.0%

Throughput +1.47% geomean, TTFT -5.02%. SCLK 2163-2400 MHz across both arms, no throttling.

The cold numbers are the ones that characterise the kernel change; the warm numbers show what a workload with high prefix reuse sees, where the prefill path is mostly skipped and the remaining TTFT gain comes from the chunks that still miss.

Across the three context lengths the gain tracks prefill share, which is the expected direction for a prefill-only change:

ISL/OSL throughput delta
1k/1k +1.4%
8k/1k +5.4%
70k/300 (cold) +13.8%

@akao-amd

Copy link
Copy Markdown
Contributor

Verified the performance uplift. LGTM.

@sogalin sogalin added the run-ci CI: run the baseline test suite on this PR label Oct 2, 2026
@ankith117

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

2 similar comments
@ankith117

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

@ankith117

Copy link
Copy Markdown
Contributor Author

/rerun-failed-ci

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

jit-kernel run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants