Repository navigation
perf(kernel): read dsv41 indexer key rows as int32 words - #634
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 slowdown of the DeepSeek-V4.1 indexer (
_indexer_logits_packed_kernel, up to +102% on H100).Call path:
Indexer(DeepSeek-V4.1, prefill and decode) ->DSV41SparseAttnBackend.indexer_logits->indexer_logits_packed->_indexer_logits_packed_kernel. Full mode runs on the Full indexer layers, candidate mode on the Reindex layers.On Triton 3.8 (sm_90) the [64 x 128] packed-row gather tile compiles with its threads along the channel axis (
threadsPerWarp [1, 32], order [1, 0]) instead of along the rows ([32, 1], order [0, 1]) as on 3.6. Registers go from 165 to 200 and the loop from 2735 to 3981 SASS instructions.Changes:
fp4_e8m0_b32, head_dim 128:IDX_FMTand DeepSeek-V4.1'sindex_head_dim),_score_tileloads the key rows with a new_load_fp4_e8m0_rows: each row is read as 17 int32 words (16 words of e2m1 codes, 1 word of ue8m0 scales), the nibbles are unpacked in registers withtl.join, e2m1 is decoded with integer bit operations and an exact multiply by 2^126, and the scale is computed once per word. Other formats and head dims keepload_rows.The gather tile now compiles to the same layout under both Triton versions; on sm_90 the kernel uses 128 registers and 1208 (3.6) / 1065 (3.8) SASS instructions in the loop.
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.1-Flash shapes: 32 index heads, head_dim 128.Over all 13 shapes, relative to old@3.6:
All shapes, H100
All shapes, RTX PRO 6000
Correctness and tests
mul.f32(checked in the PTX). Triton's AMD backend setsdenormal-fp-math-f32="ieee"by default in 3.6 and 3.8; this was not run on ROCm.tests/kernels/test_dsv41_indexer.py,test_dsv41_sparse_attn.py,test_dsv41_topk.py,test_dsv41_pack.py,tests/attention/test_dsv41_backend.py,tests/kvcache/test_dsv41_pool.pyandtests/models/deepseek_v41: 94 passed under both versions on H100 and on the RTX PRO 6000. The model tests useindex_head_dim32 and run the unchanged path;test_dsv41_indexer.pyandtest_dsv41_backend.pyrun head_dim 128.Not covered: end to end (no DeepSeek-V4.1 weights were available), timing on other GPUs (sm_80, sm_86, sm_89 and sm_100 were only compiled), ROCm.