Repository navigation
Conversation
_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.
There was a problem hiding this comment.
@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 |
There was a problem hiding this comment.
this is gfx use only triton kernel, I think you meant non-gfx1250 vs. non-gfx95 NVIDIA
There was a problem hiding this comment.
Yes. Thank you for pointing it out. Fixed it.
|
@ankith117 lint to fix. |
|
@HaiShaw lint fixed. Thank you. |
|
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.
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 |
|
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 Cache-cold (run 1 — the pass that actually prefills), n=1 per point:
Throughput +13.77% geomean, TTFT -16.66%. Cache-warm (runs 2-5, the standard median-of-4 protocol):
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:
|
|
Verified the performance uplift. LGTM. |
|
/rerun-failed-ci |
2 similar comments
|
/rerun-failed-ci |
|
/rerun-failed-ci |
Motivation
The Triton path in
_sparse_attn_v4_paged_prefill_tritoncurrently launches with a conservative configuration:block_h = 16, the AMD MFMA minimum tile, and nonum_stagesargument. 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.loadgathers throughkv_indices, which benefits from software pipelining, and withindex_n_heads = 64a 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=2on gfx1250.BLOCK_H=64equalsindex_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:
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 autotunekey=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_BUCKETkey — 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} soBLOCK_H=64exercises 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:
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
Latency, TTFT -8.1% geomean, TPOT -5.4%
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