Repository navigation
perf(kernel): keep 16-byte kv loads in dsv4 sparse attention and unroll gated pool - #633
Merged
Merged
Conversation
9 of 18 tasks
pretor
pushed a commit
to pretor/FreeToken
that referenced
this pull request
Oct 8, 2026
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.
Part of #629: fixes the Triton 3.8 slowdowns of DeepSeek-V4 sparse attention and of the compressor's gated pool.
Call paths:
DSV4SparseAttnBackend.attend->sparse_attn_paged->_sparse_attn_paged_kernel(prefill) or_sparse_attn_paged_splitk_kernel(decode)Compressor.decode_step(DeepSeek-V4) andCompressor.forward_decode(DeepSeek-V4.1) ->gated_pool->_gated_pool_kernelChanges:
sparse_attn_paged: the per-column pool basetl.where(is_win, win_ptr, cmp_ptr)is wrapped intl.multiple_of(..., 16)in both kernels, and the wrapper asserts that both pools are 16-byte aligned. On Triton 3.8 the KV gather otherwise compiles to 2-byteld.globalloads (64 in the PTX) and 82048 B of shared memory; with the hint it compiles to 16-bytecp.asyncand 67968 B under both versions. Under Triton 3.6 the new prefill kernel compiles to the same SASS as the old one on sm_80, sm_86, sm_89, sm_90, sm_100 and sm_120.gated_pool: both loops over R usetl.range(R, loop_unroll_factor=2). On Triton 3.8 the max loop issues 1 load per iteration in SASS (4 on 3.6); with the change it issues 8 under both versions.The pools passed to
sparse_attn_pagedare allocated per layer withtorch.zerosinDSV4PagedKVCache._alloc_buffers. A misaligned pool, which old code accepted, now fails the assert.Kernel time
µs,
triton.testing.do_bench_cudagraph, best of two medians, GPU used by the benchmark only. old =main377a3bc. 3.6 = torch 2.11.0 / Triton 3.6.0, 3.8 = torch 2.14.1 / Triton 3.8.0 (#630). H100 80GB HBM3 (driver 580.95.05) and RTX PRO 6000 Blackwell Server Edition (sm_120, driver 595.91.07). DeepSeek-V4 shapes: 64 heads, head_dim 512, 128-token window, index top-k 512; r = compress ratio.Over the 11 sparse attention shapes, new@3.8 is 0.94 to 1.00 of old@3.6 on H100 and 0.95 to 1.005 on the RTX PRO 6000. gated_pool at R = 8 is 0.78 to 0.81 of old@3.6 on H100 and 0.89 to 0.91 on the RTX PRO 6000, under both versions. gated_pool at R = 128 (the KV compressor of the ratio-128 layers) is 0.81 to 0.83 of old@3.6 on H100, and 1.12 to 1.13 on the RTX PRO 6000 under both versions (1.89 to 1.99 for old@3.8); three separate runs on the RTX PRO 6000 gave the same ratios.
All shapes, H100
All shapes, RTX PRO 6000
Correctness and tests
tests/kernels/test_dsv4_sparse_attn.py,tests/dsv4/test_dsv4_attn_backend.py,tests/kernels/test_dsv4_indexer_decode.py,tests/kvcache/test_dsv4_pool.py,tests/scheduler/test_dsv4_generic_manager.py, the DeepSeek-V4.1 kernel, attention and pool tests, andtests/models/deepseek_v41: 172 passed under both versions on H100 and on the RTX PRO 6000.End to end
DeepSeek-V4-Flash on H100 with the default (hybrid) MoE strategy, AIME25 problem 1 (single sampled runs; old@3.8 is
mainwith the dependency change of #630):Not covered: end to end on the RTX PRO 6000 (no weights there), timing on other GPUs (sm_80, sm_86, sm_89 and sm_100 were only compiled), ROCm. DeepSeek-V4.1 sparse attention (
dsv41/sparse_attn.py) is not changed.