Repository navigation
feat(cuda): split-KV decode attention (TLLM-ATTN-SPLITKV PR-B) - #12
Merged
Merged
Conversation
The decode attention kernel parallelises only over q-heads: at visible=2048 it
launches grid=(14,1,1) x 128 threads and achieves 8.33% occupancy with nothing
saturated (DRAM 0.44%, L2 0.96%, L1 4.0%). Split the visible KV window across
blocks instead: grid=(Hq, num_splits), each block producing a local (m, l, acc),
merged by a small combine kernel at grid=(Hq, 1).
Graph capture is preserved: split ranges are derived in-kernel from the device
visible length, num_splits is a host parameter so the grid is fixed per launch,
and the partial buffer is caller-allocated -- no allocation or D2H during
capture. Ranges are logical token ranges, so ContiguousRows/PagedRows, the
frozen address formula and the zero-row semantics are untouched.
The reduction loop is not copied: decode_online_softmax becomes
decode_online_softmax_range<Rows, kFinalize>. The two existing entry points call
it with begin=0/end=visible/kFinalize=true and are otherwise unchanged, so all
~18 existing call sites and their tests keep working.
Gates (tests/test_attention_splitkv.cpp, 11 cases):
- num_splits == 1 is bitwise identical to the single-pass path, contiguous and
paged. exp(0) is forced to exactly 1.0f for the max split rather than trusting
the fast intrinsic, which makes that exactness structural, not incidental;
- num_splits in {2,4,8,16} stays inside the independent oracle's 2e-3 tolerance;
- split contiguous and split paged agree bitwise (same reduction, two addressing
strategies);
- visible=0, visible<num_splits (empty splits), non-multiple-of-splits, illegal
block ids, short block table;
- CUDA Graph capture, then replay with a growing visible length, matches eager --
the direct proof of the graph-capture premise rather than an assertion of it;
- non-default stream matches the default stream.
Full suite 219 passed / 11 skipped, unchanged from before this commit.
Co-authored-by: CommandCodeBot <noreply@commandcode.ai>
…tic late max Mutation testing over the split-KV gates (4 injections): - combine drops the l_i weight -> caught by 6 oracle-based tests, while the num_splits = 1 anchor still passes (w == 1 there). This is the concrete demonstration that the anchor alone is not enough and the independent oracle is load-bearing. - split range off-by-one -> caught by 7 tests, including both bitwise anchors. - drop combine's `m == M -> 1.0f` special case -> not caught, and that is correct: a probe confirms __expf(0.0f) == 1.0f exactly (+0 and -0), so the special case is a no-op on this toolchain. It is kept only so the anchor does not depend on __expf's behaviour at 0. - drop the online rescale in the shared loop -> NOT caught by the original matrix. The last one is a real gap in the tests, not a gate failure: with random data and short sequences the global max almost always lands in the first tile, so old_rescale is exp(0) = 1 and the rescale path is never exercised -- the injected bug passed all 11 gates. LateMaxForcesOnlineRescaleToMatter constructs the opposite deterministically (K all zero except one token in the last tile, which therefore holds the global max), forcing old_rescale != 1; the mutation is now caught by that test alone. Corrects the design record's mutation list, which had claimed the independent oracle would catch a dropped rescale. That was only true once the matrix actually reached a multi-tile rescale, which it did not. Full suite 220 passed / 11 skipped; sanitizer 0 errors on the new tests and the existing paged suites; clang-format 18.1.8 clean. Co-authored-by: CommandCodeBot <noreply@commandcode.ai>
…DA 11.8 CI builds against CUDA 11.8 while the dev box has 13.3. The three-argument cudaGraphInstantiate overload is CUDA 12+, so PR-B failed to build there: error: too few arguments to function 'cudaError_t cudaGraphInstantiate(...)' Switch to cudaGraphInstantiateWithFlags, which exists since CUDA 11.4 and is present in both toolchains. The five-argument legacy form was the alternative but is not guaranteed to still exist in CUDA 13. Caught by CI, not by the local build: the kernel changes themselves compiled fine on 11.8, so this was purely a test-side API-version issue. Co-authored-by: CommandCodeBot <noreply@commandcode.ai>
This was referenced Sep 14, 2026
Merged
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.
TLLM-ATTN-SPLITKV PR-B:kernel + 门禁。stacked on #11(设计包)。本 PR 纯增量:新增入口点,尚未接线,生产行为不变(接线是 PR-C)。
为什么
PR-A 的诊断:decode attention 只有
num_q_heads一个并行轴,visible=2048时grid=(14,1,1)×128、achieved occupancy 8.33%,且无资源饱和(DRAM 0.44%、SM 0.72%、L2 0.96%)。修法是加 block:把可见 KV 逻辑窗口切段,grid=(Hq, num_splits),每段产出局部(m,l,acc),再由grid=(Hq,1)的 combine kernel 合并。新增行为
attention_decode_splitkv/attention_decode_paged_splitkv(内部头,C ABI 不变)。decode_online_softmax→decode_online_softmax_range<Rows, kFinalize>:不复制归约,两条路径继续共用同一份循环;既有两个入口以begin=0/end=visible/kFinalize=true调用,逐字保持原行为。ContiguousRows/PagedRows、冻结地址公式、零 row 语义全未改动。visible_tokens现场派生,num_splits是 host 参数(grid 捕获时固定),partial 缓冲由调用方预分配、捕获期零分配零 D2H。验证
num_splits == 1与单遍路径逐位相同(连续 + 分页)。m == M时权重显式取1.0f,使该性质不依赖__expf(0)。num_splits ∈ {2,4,8,16}落在独立 oracle 的2e-3容差内;split 连续 vs split 分页逐位相同。visible=0、visible < num_splits(空段)、非整数倍、非法块 id、短块表。compute-sanitizer新用例与既有 paged suites 0 error;clang-format 18.1.8 0 violation。变异检验(4 项,含一个真实缺口)
l_i权重num_splits=1锚点照常通过(那里w==1)——说明锚点独自不够,独立 oracle 是必需的m==M→1.0f__expf(0.0f)==1.0f精确成立,该特例在本工具链是 no-op最后一项是测试缺口而非门禁失效:随机数据 + 短序列时全局 max 几乎总在第一个 tile,
old_rescale恒为exp(0)=1,rescale 路径根本没被压到(注入后 11 项门禁全过)。已补LateMaxForcesOnlineRescaleToMatter:确定性构造后置 max(末 tile 单 token 取大 K、其余为 0),强制old_rescale ≠ 1,此后该变异被单独抓出。设计包的变异清单已按事实更正。限制
src/transformer.cpp调用,LayerWorkspace与开关是 PR-C。num_splits > 1会改变 fp32 求和顺序,因此不再与现状逐位相同(这正是设计包 Q3,owner 已批准);num_splits=1逐位锚点是硬门禁。num_splits的最优值与"小 S 是否回归"仍是待实测项。后续
PR-C(
LayerWorkspace+TLLM_ATTN_SPLITKV+transformer.cpp路由)→ PR-D(benchmark + 结果报告)。