Skip to content

[Kimi-K3] Add GEMM-RS for sequence parallelism - #52079

Merged
simon-mo merged 4 commits into
vllm-project:mainfrom
gau-nernst:codex/kimi-k3-gemm-rs
Aug 13, 2026
Merged

simon-mo merged 4 commits into
vllm-project:mainfrom
gau-nernst:codex/kimi-k3-gemm-rs

Conversation

@gau-nernst

@gau-nernst gau-nernst commented Aug 13, 2026 •

Copy link
Copy Markdown
Contributor

Purpose

Add GEMM-RS kernel for Blackwell, based on https://github.com/NVIDIA/cutlass/blob/dcf215a/examples/python/CuTeDSL/cute/blackwell/kernel/distributed/distributed_gemm_reduce_scatter_blackwell.py (multimem.ld_reduce)

  • Supports any value of M e.g. M=1023. However, only uses GEMM-RS when M>=128 since the kernel was not optimized for small/medium M. Only supports TP<=16, and requires all rank on the same NVLink domain.
  • The sharding behavior follows existing RS logic i.e. eac rank holds ceil(M / world_size), with the exception of the last rank
  • Requires an opt-in flag VLLM_KIMI_K3_GEMM_RS, which enables GEMM-RS for O-proj and shared experts+dense MLP down-proj
  • Symmetric memory workspace (per-GPU): max_num_batched_tokens x 7168 x 2 bytes = 448 MiB for MNBT=32k. Confirmed in vLLM logs KV memory 39.99 GiB (before) -> 39.78 GiB (after) -> not much

Initialization and runtime logic

  • Whether to initialize GEMM-RS: done in maybe_init_gemm_rs(), which also logs the reason if it fails. When VLLM_KIMI_K3_GEMM_RS=0, it doesn't do anything
  • At each layer's __init__(), we call self.run_gemm_rs = get_gemm_rs().can_run(self.down_proj.weight). This is to further validate supported weight shapes and dtype
  • When the checks fail, we fallback to standard behavior.
  • In forward(), we check again with should_run(), which is the heuristics M>=128. The kernel supports any values of M, but right now the baseline is better for M<128

Though technically this can work with any SP in general, this PR only enables GEMM-RS for Kimi-K3. A future extension is to make this into GEMM-AR by adding multimem.st (all-gather) after multimem.ld_reduce (reduce-scatter).

Microbenchmark

benchmarks/kernels/benchmark_kimi_k3_gemm_rs.py in this PR. CUDA graph with rotating buffers. All benchmarks were done with GB300.

TP4

Note: K=1536 is shared expert down-proj, K=3072 is O-proj

M N K Torch GEMM + NCCL RS (RING_LL) (us) Torch GEMM + NCCL RS (LDMC) (us) GEMM-RS (us) Speedup vs RING_LL Speedup vs LDMC
128 7168 1536 48.13 46.27 38.74 1.242 1.195
512 7168 1536 84.64 60.51 50.69 1.67 1.194
2048 7168 1536 119.92 110.83 83.74 1.432 1.323
8192 7168 1536 286.58 343.22 209.25 1.37 1.64
32768 7168 1536 1044.66 1265.86 719.7 1.452 1.759
128 7168 3072 52.37 48.27 43.38 1.207 1.113
512 7168 3072 92.35 66 53.09 1.74 1.243
2048 7168 3072 145.6 133.81 89.44 1.628 1.496
8192 7168 3072 383.65 431.82 234.38 1.637 1.842
32768 7168 3072 1402.05 1628.43 1001.49 1.4 1.626

Component breakdown

M N K Torch GEMM (us) NCCL RS (best) (us) GEMM-RS (us)
128 7168 1536 17.65 39.89 38.74
512 7168 1536 21.52 50.94 50.69
2048 7168 1536 36.54 90.37 83.74
8192 7168 1536 104.14 191.68 209.25
32768 7168 1536 401.31 652.77 719.7
128 7168 3072 18.96 41.38 43.38
512 7168 3072 28.21 50.48 53.09
2048 7168 3072 59.25 90.56 89.44
8192 7168 3072 195.92 196.91 234.38
32768 7168 3072 763.71 650.18 1001.49

TP8

Note: K=768 is shared expert down-proj, K=1536 is O-proj

M N K Torch GEMM + NCCL RS (RING_LL) (us) Torch GEMM + NCCL RS (LDMC) (us) GEMM-RS (us) Speedup vs RING_LL Speedup vs LDMC
128 7168 768 49.47 44.93 40.08 1.234 1.121
512 7168 768 70.29 56.58 47.46 1.481 1.192
2048 7168 768 110.34 100.75 81.1 1.36 1.242
8192 7168 768 264.43 300.98 202.27 1.307 1.488
32768 7168 768 928.83 1109.26 697.5 1.332 1.59
128 7168 1536 51.26 46.86 41.63 1.231 1.126
512 7168 1536 72.22 62.64 49.82 1.45 1.257
2048 7168 1536 121.71 109.84 82.08 1.483 1.338
8192 7168 1536 306.91 347.92 209.04 1.468 1.664
32768 7168 1536 1108.27 1289.57 706.02 1.57 1.827

Component breakdown

M N K Torch GEMM (us) NCCL RS (best) (us) GEMM-RS (us)
128 7168 768 17.3 39.2 40.08
512 7168 768 18.61 49.44 47.46
2048 7168 768 26.64 91.39 81.1
8192 7168 768 60.18 209.14 202.27
32768 7168 768 216.29 714.27 697.5
128 7168 1536 17.22 39.66 41.63
512 7168 1536 23.17 51.39 49.82
2048 7168 1536 36.64 90.43 82.08
8192 7168 1536 104.14 210.16 209.04
32768 7168 1536 400.02 713.71 706.02

E2E prefill-only benchmark

All benchmarks were done with 8xGB300, TP8+EP+SP (DeepGEMM MegaMoE), --max-num-batched-tokens 32768, 8k input - 1 output requests. Baseline is 7aa248f

Concurrency Baseline TTFT (median) GEMM-RS TTFT (median) Baseline TPGS GEMM-RS TPGS
C1 312.81 ms 298.44 ms (-4.59%) 3,172.7 tok/GPU/s 3,432.2 tok/GPU/s (+8.18%)
C32 7,240.71 ms 6,799.51 ms (-6.09%) 4,500.5 tok/GPU/s 4,794.0 tok/GPU/s (+6.52%)

Test Plan

Unit test (also added to distributed CI)

tests/kernels/test_kimi_k3_gemm_rs.py

E2E testing, TP8+EP+SP (DeepGEMM MegaMoE)

  • GSM8K: 96.82%
  • OCRBench: 88.40%

Test Result


Essential Elements of an Effective PR Description Checklist
  • The purpose of the PR, such as "Fix some issue (link existing issues this PR will resolve)".
  • The test plan, such as providing test command.
  • The test results, such as pasting the results comparison before and after, or e2e results
  • (Optional) The necessary documentation update, such as updating supported_models.md and examples for a new model.

BEFORE SUBMITTING, PLEASE READ https://docs.vllm.ai/en/latest/contributing (anything written below this line will be removed by GitHub Actions)

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Claude Code Review

This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.

@mergify mergify Bot added ci/build performance Performance-related issues kimi k3 labels Aug 13, 2026
Comment thread vllm/models/kimi_k3/nvidia/ops/cute_dsl/gemm_rs.py
@gau-nernst

Copy link
Copy Markdown
Contributor Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #83638 for commit 1f6da41685b6.

Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
@gau-nernst
gau-nernst force-pushed the codex/kimi-k3-gemm-rs branch from 4210d0f to 50bc8b7 Compare August 13, 2026 01:52
@gau-nernst

Copy link
Copy Markdown
Contributor Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #83645 for commit 50bc8b7ba072.

@gau-nernst

Copy link
Copy Markdown
Contributor Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #83667 for commit eb8c5ffc1391.

@gau-nernst

Copy link
Copy Markdown
Contributor Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #83689 for commit a47948a94ac5.

@simon-mo simon-mo 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.

stamp

@simon-mo
simon-mo merged commit 6014f9e into vllm-project:main Aug 13, 2026
97 checks passed
vrdn-23 added a commit to vrdn-23/vllm that referenced this pull request Aug 14, 2026
Resolves the recurring vllm/envs.py structural conflict per
docs/superpowers/specs/2026-05-14-envs-merge-conflict-resolution-design.md:
main's legacy `if TYPE_CHECKING:` block and `environment_variables` dict are
dropped wholesale (superseded by the pydantic BaseSettings tree on this
branch), then main's semantic delta is ported field-by-field.

6 main-side commits touched vllm/envs.py since base e644c8c (+55 -0).
All 10 new vars have already-merged callers, so every port is mandatory:

- vllm-project#51447 VLLM_MAX_STOP_STRINGS (int=4), VLLM_MAX_NUM_BAD_WORDS (int=128),
  VLLM_MAX_BAD_WORDS_TOTAL_TOKENS (int=1024) -> ServerSettings
- vllm-project#49948 VLLM_MAX_AUDIO_DECODE_BYTES (int=268_435_456) -> MediaSettings,
  carrying compile_factor=False to mirror main's ignore-set addition
- vllm-project#50484 VLLM_USE_DIRECT_DCP_A2A / _Q_GATHER / _KV_GATHER (bool|None=None)
  -> QuantSettings, with one shared `_parse_direct_dcp` before-validator
  reproducing main's maybe_convert_bool exactly
- vllm-project#52079 VLLM_KIMI_K3_GEMM_RS (bool=False) -> QuantSettings
- vllm-project#49458 VLLM_USE_HW_AGNOSTIC (bool=False) -> UsageSettings
- vllm-project#47808 VLLM_ADAPTIVE_VERIFICATION_PROFILE_CONTEXT_LEN (int=8192)
  -> QuantSettings

No deletions, modifications, or renames this window. Nothing was dropped
silently: all 6 commits' envs.py deltas are covered above.

Env var set parity after resolution: 292 branch fields vs 293 main runtime
entries, sole difference VLLM_TRITON_ATTN_USE_TD -- the known deprecation
shim divergence, re-confirmed untouched by this merge window.

Verified: 54 tests pass across tests/test_envs.py, tests/test_envs_pydantic.py
and tests/docs/test_env_vars_gen.py; `pre-commit run --files vllm/envs.py`
clean; tests/test_request_input_bounds.py passes (22 tests). The audio, DCP
and end-to-end hw-agnostic consumer suites need a CUDA box plus soundfile /
multiprocess and were not run here.

AI assistance was used to enumerate the port list and apply the resolution;
see Appendix G of the playbook for the full audit trail.

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Vinay Damodaran <vrdn@hey.com>
@gau-nernst
gau-nernst deleted the codex/kimi-k3-gemm-rs branch August 14, 2026 01:27
zufangzhu pushed a commit to zufangzhu/vllm that referenced this pull request Aug 24, 2026
Signed-off-by: Thien Tran <gau.nernst@yahoo.com.sg>
Signed-off-by: Zhu, Zufang <zufang.zhu@intel.com>

ptr = x.iterator.toint(loc=loc, ip=ip).ir_value(loc=loc, ip=ip)
asm = (
"multimem.ld_reduce.relaxed.gpu.global.add.acc::f32"

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

.sys scope for multimem.ld_reduce

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

functionally speaking, .gpu scope also work on hw. PTX semantic speaking, .sys scope is better.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

ci/build k3 kimi performance Performance-related issues

3 participants