Skip to content

[Kimi-K3] Enable KDA projection fusion for DP attention - #36206

Closed
ajit283 wants to merge 6 commits into
sgl-project:mainfrom
ajit283:kimi-k3-dep16-fused-proj
Closed

ajit283 wants to merge 6 commits into
sgl-project:mainfrom
ajit283:kimi-k3-dep16-fused-proj

Conversation

@ajit283

@ajit283 ajit283 commented Aug 24, 2026 •

Copy link
Copy Markdown

Motivation

Kimi-K3's full-rank KDA projection fusion was disabled whenever attention TP differed from global TP. That excludes DP-attention deployments such as TP16 + DP-attention16, even though the full-rank merged QKVG projection accepts explicit attention-TP rank/size and can safely use the replicated attention weights.

This enables the existing fusion for that layout. No new GPU kernel is introduced.

Modifications

  • Allow full-rank KDA projection fusion when attn_tp_size != tp_size, including mixed block-FP8 attention projections.
  • Compute fused QKVG output splits using attn_tp_size.
  • Preserve the existing restriction for the low-rank repeated/batched layout, which still requires matching attention/global TP and no quantization config.
  • Add unit coverage for the full-rank DP-attention and low-rank policies, direct fused/unfused projection math, and CUDA-graph side-stream replay.
  • Add a CUDA-graph projection microbenchmark and an Nsight slice comparison utility under benchmark/kernels/kimi_k3/.

For DEP16, this changes the projection prologue from five logical GEMMs:

QKV + beta + f_a + f_b + gate

to three:

QKVG + BFA + f_b

Accuracy Tests

Tested with the real Kimi-K3 checkpoint on 16 GB300 GPUs, TP16 + DP-attention16 + EP16:

  • GPU unit job 630507: 7 passed, including five parameterized subtests.
    • Fused projections match the corresponding unfused BF16 math at M={1,32}.
    • Side-stream eager and CUDA-graph replay outputs match the serial path.
  • Deterministic full-model A/B job 631695:
    • 32 fixed prompts, 64 generated tokens each.
    • 32/32 generated texts identical.
    • 2048/2048 output token IDs identical.
    • 2048/2048 output token logprobs identical; maximum absolute difference 0.0.
  • Candidate job 523990 loaded the real checkpoint, captured all CUDA graphs, and served 2048/2048 requests.
  • Post-rebase pre-commit, Python compilation, and git diff --check pass.

A preliminary normal-mode comparison produced different outputs, but a baseline-versus-baseline control also diverged despite fixed server seeds and sequential requests. The normal serving configuration is not batch-invariant and therefore is not a valid exact-equality oracle. Enabling SGLang deterministic inference made the baseline/candidate comparison bit-identical as reported above.

GSM8K

Full 1,319-question GSM8K test split, standard 5-shot completion evaluation, greedy decoding, and three independent server starts per variant:

Variant Scores Mean Median
Unfused 95.814%, 95.358%, 95.510% 95.561% 95.510%
Fused 95.282%, 95.586%, 94.673% 95.180% 95.282%

The fused-minus-unfused difference is -0.381 percentage points by mean and -0.228 points by median, within the observed normal-mode run-to-run spread. The deterministic full-model comparison above remains bit-identical.

Speed Tests and Profiling

Measured on NVIDIA GB300 (SM103), TP16 + DP-attention16 + EP16, global concurrency 512 (local decode M=32), ISL/OSL 8192/1024, FP8 KV cache, BF16 SSM state, and CUDA graphs.

Repeated unprofiled serving A/B

Job 630711, three repetitions per variant on the same four-node NVL72 allocation:

Metric Unfused median Fused median Change
Output throughput 3055.28 tok/s 3122.06 tok/s +2.19%
Request throughput 2.9837 req/s 3.0489 req/s +2.19%
Mean TPOT 95.160 ms 93.123 ms -2.14%
Mean E2E latency 161.229 s 157.630 s -2.23%

Per-run output throughput:

Unfused: 3038.33, 3055.28, 3061.53 tok/s
Fused:   3106.03, 3133.26, 3122.06 tok/s

Nsight profiling

Metric Unfused Fused Change
Individual KDA projection prologue 140.064 us 117.281 us -16.27%
Three projection prologues in aligned slice 420.304 us 351.697 us -16.32%
Aligned four-layer wall time 1977.695 us 1951.731 us -1.31%
Full graph median 45.262 ms 44.807 ms -1.00%
Kernels per four-layer slice 114 105 -7.89%
CUDA graph nodes 2651 2444 -7.81%

Summed kernel duration increases from 2190.766 us to 2363.971 us because the side-stream BFA/f_b work overlaps QKVG and runs longer under SM contention. Wall time and full-graph latency improve, so those are the relevant end-to-end metrics.

The branch was subsequently rebased to current main; the tested kimi_k3.py remained byte-identical across that rebase.

Checklist


CI States

Latest PR Test (Base): ⏳ Run #35610810492
Latest PR Test (Extra): ❌ Run #35610809019
Latest PR Test (AMD ROCm 10): ⏳ Run #35610809647

@github-actions github-actions Bot added the documentation Improvements or additions to documentation label Aug 24, 2026
@ajit283
ajit283 force-pushed the kimi-k3-dep16-fused-proj branch 4 times, most recently from bcca1af to 5e4eee3 Compare September 1, 2026 06:24
@ajit283
ajit283 marked this pull request as ready for review September 1, 2026 06:25
@nvpohanh

nvpohanh commented Sep 1, 2026

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Sep 1, 2026
@ajit283
ajit283 force-pushed the kimi-k3-dep16-fused-proj branch from 71d1001 to 0565053 Compare September 1, 2026 14:50
@nvpohanh

nvpohanh commented Sep 7, 2026

Copy link
Copy Markdown
Collaborator

/rerun-failed-ci

1 similar comment
@nvpohanh

nvpohanh commented Sep 8, 2026

Copy link
Copy Markdown
Collaborator

/rerun-failed-ci

@nvpohanh

nvpohanh commented Sep 8, 2026

Copy link
Copy Markdown
Collaborator

/rerun-failed-ci

@nvpohanh

nvpohanh commented Sep 9, 2026

Copy link
Copy Markdown
Collaborator

All NV pipelines have passed

@nvpohanh

Copy link
Copy Markdown
Collaborator

cc @DarkSharpness since this is related to K3

@kpham-sgl kpham-sgl left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can you also test with GSM8K and/or AIME26? Thanks!

Comment thread benchmark/kernels/kimi_k3/bench_kda_dep16_projections.py Outdated
Comment thread python/sglang/srt/models/kimi_k3.py Outdated
@nvpohanh

Copy link
Copy Markdown
Collaborator

@ajit283 please address the comments. Thanks!

@ajit283

ajit283 commented Sep 18, 2026

Copy link
Copy Markdown
Author

@kpham-sgl I ran the full 1,319-question GSM8K eval three times per variant. Baseline averaged 95.56% and this PR averaged 95.18% (medians 95.51% vs 95.28%). That is a -0.38 pp mean / -0.23 pp median difference, within the run-to-run spread we saw in normal serving. The deterministic full-model A/B is still bit-identical. I added the full numbers to the PR description.

@nvpohanh

Copy link
Copy Markdown
Collaborator

@ajit283 please fix the conflicts. thanks

…-proj

# Conflicts:
#	python/sglang/srt/models/kimi_k3.py
@kpham-sgl

Copy link
Copy Markdown
Collaborator

@ajit283 @nvpohanh seems like #39589 already landed most of this PR's functionality.... We can close this one out

@kpham-sgl kpham-sgl closed this Sep 21, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants