Skip to content

feat(cuda): split-KV decode attention (TLLM-ATTN-SPLITKV PR-B) - #12

Merged
holtwood merged 3 commits into
masterfrom
tllm-attn-splitkv-kernel
Sep 15, 2026
Merged

holtwood merged 3 commits into
masterfrom
tllm-attn-splitkv-kernel

Conversation

@holtwood

Copy link
Copy Markdown
Member

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 调用,逐字保持原行为。
  • 段边界落在逻辑 token 上 ⇒ ContiguousRows/PagedRows、冻结地址公式、零 row 语义全未改动。
  • 保住 CUDA Graph:段范围由 device 端 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、短块表。
  • CUDA Graph:捕获后让可见长度从 1 增长到 96 并 replay,与 eager 逐位相同——直接验证"grid 固定 + device int"这一前提。
  • 全量 220 passed / 11 skipped(基线 208,新增 12);compute-sanitizer 新用例与既有 paged suites 0 error;clang-format 18.1.8 0 violation。

变异检验(4 项,含一个真实缺口)

注入 结果
combine 丢掉 l_i 权重 被 6 项 oracle 用例抓到;num_splits=1 锚点照常通过(那里 w==1)——说明锚点独自不够,独立 oracle 是必需的
段范围 off-by-one 被 7 项抓到,含两条锚点
去掉 combine 的 m==M→1.0f 未抓到且应为未抓到:探针证实 __expf(0.0f)==1.0f 精确成立,该特例在本工具链是 no-op
丢掉共享循环的在线 rescale 初版矩阵漏掉

最后一项是测试缺口而非门禁失效:随机数据 + 短序列时全局 max 几乎总在第一个 tile,old_rescale 恒为 exp(0)=1,rescale 路径根本没被压到(注入后 11 项门禁全过)。已补 LateMaxForcesOnlineRescaleToMatter:确定性构造后置 max(末 tile 单 token 取大 K、其余为 0),强制 old_rescale ≠ 1,此后该变异被单独抓出。设计包的变异清单已按事实更正。

限制

  • 本 PR 不改变生产路径:两个新入口还没有被 src/transformer.cpp 调用,LayerWorkspace 与开关是 PR-C。
  • num_splits > 1 会改变 fp32 求和顺序,因此不再与现状逐位相同(这正是设计包 Q3,owner 已批准);num_splits=1 逐位锚点是硬门禁。
  • 不产生任何性能数字,也不声称收益——那是 PR-D。num_splits 的最优值与"小 S 是否回归"仍是待实测项。
  • 作者即实现者,§12 的独立性缺陷仍成立:合并前需要第二方复核 diff。

后续

PR-C(LayerWorkspace + TLLM_ATTN_SPLITKV + transformer.cpp 路由)→ PR-D(benchmark + 结果报告)。

holtwood and others added 3 commits September 14, 2026 22:30
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>
@holtwood
holtwood changed the base branch from tllm-attn-splitkv-design to master September 15, 2026 04:21
@holtwood
holtwood merged commit b829ba8 into master Sep 15, 2026
2 checks passed
@holtwood
holtwood deleted the tllm-attn-splitkv-kernel branch September 16, 2026 06:19
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant