From 340cc0202d13458f8ae2953f58df4e1bf89f9c47 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:32:24 -0700 Subject: [PATCH 01/35] docs: design Kimi-Linear CP-v2 transitions --- .../2026-07-17-kimi-linear-cp-v2-design.md | 134 ++++++++++++++++++ 1 file changed, 134 insertions(+) create mode 100644 docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md diff --git a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md new file mode 100644 index 000000000000..830ac41bfde7 --- /dev/null +++ b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md @@ -0,0 +1,134 @@ +# Kimi-Linear CP-v2 Layer Transitions + +## Context + +PR #31619 moves MLA prefill context parallelism to the CP-v2 strategy API. At +the model boundary, CP-v2 splits embeddings and positions into the selected CP +layout, and after the model body it gathers hidden states before logits. + +Kimi-Linear alternates two attention implementations with different token +layout requirements: + +- Kimi Delta Attention (KDA) is tensor parallel and must receive the complete + token batch on every TP rank. +- Multi-head Latent Attention (MLA) participates in prefill context parallelism + and must receive the rank-local zigzag token shard. + +The model carries both `hidden_states` and a deferred `residual` between decoder +layers. They must always use the same token layout. + +## Goals + +- Support Kimi-Linear prefill with `--tp 4 --attn-cp-size 4 + --enable-prefill-cp --cp-strategy zigzag`. +- Gather CP-sharded layer state before each KDA region. +- Split replicated layer state before each MLA region. +- Reuse the active CP strategy for ordering, padding, and collectives. +- Preserve the non-CP and decode paths. +- Verify correctness with focused unit tests and GSM8K on four GB300 GPUs. + +## Non-goals + +- Adding a new CP strategy or attention backend. +- Enabling the unfinished interleave strategy. +- Changing KDA kernels, MLA kernels, or KV-cache semantics. +- Supporting Kimi K2.5's multimodal wrapper in this change. +- Optimizing transition collectives beyond the minimum correct implementation. + +## Design + +### Communicator + +Add `KimiLinearCPV2LayerCommunicator` under +`python/sglang/srt/layers/cp/kimi_linear.py`. The class owns only the transition +between the two token layouts; it does not replace the general TP/DP/MoE +`LayerCommunicator` in `layers/communicator.py`. + +Each `KimiDecoderLayer` constructs the communicator with: + +- Whether the current layer is KDA. +- Whether the preceding global layer is KDA. Layer zero has no preceding layer. + +At the beginning of `KimiDecoderLayer.forward`, the communicator receives +`hidden_states`, `residual`, and `forward_batch`, and returns state in the layout +required by the current layer. + +It is active only when `is_cp_v2_active(forward_batch)` is true. Otherwise it +returns its inputs without communication. + +### Transition table + +| Incoming state | Current layer | Operation | +| --- | --- | --- | +| Model-entry CP shard | First KDA | Gather to complete token order | +| Replicated KDA output | KDA | No-op | +| Replicated KDA output | MLA | Split with the active CP strategy | +| CP-sharded MLA output | MLA | No-op | +| CP-sharded MLA output | KDA | Gather to complete token order | + +The first Kimi-Linear layer is KDA, so model-entry embeddings are gathered +before layer zero. The final Kimi-Linear layer is MLA, so it remains CP-sharded; +the CP-v2 eager runner performs the existing model-exit gather before logits. + +### Hidden state and residual + +When `residual` is present, both tensors undergo the same transition. A gather +uses `ContextParallelStrategy.gather_hidden_states` so zigzag rank ordering, +ragged batches, and padding are restored exactly as at the model boundary. A +split uses `ContextParallelStrategy.shard_hidden_states`. + +The initial layer has `residual=None`; only `hidden_states` is gathered there. +No residual addition is moved across the transition. + +### Positions + +Position IDs remain CP-sharded after the model-entry split. KDA's forward path +does not consume `positions`, while MLA requires positions aligned with its +rank-local hidden-state shard. Consequently, position IDs do not need gather +and split transitions. + +### CP-v2 activation + +Add `KimiLinearForCausalLM` to `CP_V2_DEFAULT_MODEL_CLASSES`. Also expose the +model's input embedding layer through `get_input_embeddings`, which the CP-v2 +eager runner uses to embed the complete token batch before its initial split. + +The change remains restricted to CP-v2 context-parallel extend batches through +the existing `is_cp_v2_active` gate. Decode and legacy CP-v1 behavior are +unchanged. + +## Error handling and invariants + +- The communicator requires an initialized CP strategy whenever CP-v2 is + active; this follows the existing CP-v2 invariant. +- Both tensors must have matching token dimensions before a joint transition. +- The strategy metadata prepared by the eager runner is the single source of + truth for split and gather ordering. +- A transition is selected from static model layer types, not inferred from + runtime tensor lengths. + +## Testing + +Focused CPU unit tests will use a recording strategy to verify: + +- Model entry to first KDA gathers `hidden_states` and accepts `residual=None`. +- KDA-to-MLA splits both `hidden_states` and `residual`. +- MLA-to-KDA gathers both tensors. +- KDA-to-KDA, MLA-to-MLA, decode, and inactive CP-v2 paths are no-ops. +- Communicator integration uses the model's configured KDA/MLA layer sequence. + +The existing zigzag strategy tests cover permutation and ragged-batch ordering; +an additional transition round-trip test will ensure the communicator uses +those strategy operations without changing order. + +End-to-end verification will launch Kimi-Linear on `baizhou-dev-2` with four +GB300 GPUs and the requested flags, then run the repository's GSM8K evaluation. +The result will be compared with the same model's non-CP baseline or its known +expected accuracy, and server logs will be checked for collective, shape, and +KV-cache errors. + +## Delivery + +The implementation will be a draft stacked PR targeting +`sgl-project/sglang:cp-v2-mla-prefill`. After PR #31619 merges, the PR can be +retargeted or rebased onto `main`. From f3f6bf4641c6f9f53d96d3ad23a3778d8fb18134 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:34:54 -0700 Subject: [PATCH 02/35] docs: plan Kimi-Linear CP-v2 implementation --- .../plans/2026-07-17-kimi-linear-cp-v2.md | 248 ++++++++++++++++++ 1 file changed, 248 insertions(+) create mode 100644 docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md diff --git a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md new file mode 100644 index 000000000000..171c3094b910 --- /dev/null +++ b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md @@ -0,0 +1,248 @@ +# Kimi-Linear CP-v2 Implementation Plan + +> **For Codex:** Execute this plan test-first in the current branch. Keep the +> production change limited to the Kimi-Linear layer transitions and CP-v2 +> activation hooks. + +**Goal:** Run Kimi-Linear MLA layers on zigzag CP shards while running KDA layers +on complete token batches replicated across the four TP ranks. + +**Architecture:** A Kimi-specific CP-v2 layer communicator converts +`hidden_states` and `residual` at KDA/MLA boundaries using the active +`ContextParallelStrategy`. `KimiDecoderLayer` invokes it before input RMSNorm. +The existing CP-v2 eager runner continues to own model-entry splitting and +model-exit gathering. + +**Tech Stack:** Python, PyTorch distributed collectives, SGLang CP-v2 strategy +API, `unittest`, four NVIDIA GB300 GPUs, GSM8K evaluation. + +--- + +### Task 1: Specify layer-transition behavior with a failing unit test + +**Files:** + +- Create: `test/registered/cp/test_kimi_linear_cp_v2.py` +- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` + +**Step 1: Write the first failing test** + +Create a CPU-registered `CustomTestCase` with a recording CP strategy. The first +test constructs a communicator for the first KDA layer, patches CP-v2 active, +and asserts that its rank-local `hidden_states` are passed through +`gather_hidden_states` while `residual=None` is preserved. + +```python +communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=True, + previous_is_kda_layer=None, +) +hidden_states, residual = communicator.prepare_attn( + hidden_states, None, forward_batch +) +self.assertEqual(strategy.gather_calls, 1) +self.assertIsNone(residual) +``` + +**Step 2: Run the focused test and confirm RED** + +Run: + +```bash +python -m unittest test.registered.cp.test_kimi_linear_cp_v2 -v +``` + +Expected: failure because the communicator module/class does not exist. + +### Task 2: Implement the minimal CP-v2 communicator + +**Files:** + +- Create: `python/sglang/srt/layers/cp/kimi_linear.py` +- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` +- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` + +**Step 1: Implement first-KDA gather** + +Define `KimiLinearCPV2LayerCommunicator` with static layer-type inputs and a +`prepare_attn` method. Gate it with `is_cp_v2_active`; obtain the strategy with +`get_cp_strategy`; gather the model-entry shard for a first KDA layer. + +**Step 2: Run the focused test and confirm GREEN** + +Run the same unit-test command and require it to pass. + +**Step 3: Add one failing transition test at a time** + +Add and run tests for: + +- KDA to MLA: shard `hidden_states` and `residual`. +- MLA to KDA: gather `hidden_states` and `residual`. +- KDA to KDA and MLA to MLA: identity/no strategy calls. +- CP-v2 inactive: identity/no strategy calls. + +For each case, first observe failure, then implement the smallest transition +logic needed to pass it. + +**Step 4: Add a real zigzag round-trip test** + +Use `ZigzagCPStrategy` metadata and the existing fake CP group pattern to prove +that a full tensor split for MLA and gathered for KDA returns to original token +order, including its residual tensor. + +**Step 5: Run communicator and existing strategy tests** + +```bash +python -m unittest \ + test.registered.cp.test_kimi_linear_cp_v2 \ + test.registered.cp.test_cp_strategy_unit -v +``` + +Expected: all pass. + +### Task 3: Wire the communicator into Kimi-Linear + +**Files:** + +- Modify: `python/sglang/srt/models/kimi_linear.py` +- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` +- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` + +**Step 1: Add a failing wiring test** + +Verify the model exposes a communicator configured from each layer's current +and previous KDA/MLA type, including `previous_is_kda_layer=None` for layer zero. +Use a minimal Kimi config or constructor patching so the test remains CPU-only. + +**Step 2: Construct and invoke the communicator** + +In `KimiDecoderLayer.__init__`, compute current and previous layer types from +`KimiLinearConfig.is_kda_layer`. Construct the communicator. At the very start +of `forward`, call `prepare_attn` before input RMSNorm. + +**Step 3: Run the focused test and confirm GREEN** + +```bash +python -m unittest test.registered.cp.test_kimi_linear_cp_v2 -v +``` + +### Task 4: Enable Kimi-Linear in the CP-v2 eager path + +**Files:** + +- Modify: `python/sglang/srt/layers/cp/utils.py` +- Modify: `python/sglang/srt/models/kimi_linear.py` +- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` + +**Step 1: Add failing activation/accessor assertions** + +Assert that `KimiLinearForCausalLM` is in `CP_V2_DEFAULT_MODEL_CLASSES` and +that its `get_input_embeddings()` accessor returns `model.embed_tokens`. + +**Step 2: Add the activation and embedding hooks** + +Add the architecture string to the default class set and the standard accessor +to `KimiLinearForCausalLM`. + +**Step 3: Run the focused CP tests** + +```bash +python -m unittest \ + test.registered.cp.test_kimi_linear_cp_v2 \ + test.registered.cp.test_cp_strategy_unit -v +``` + +Expected: all pass. + +### Task 5: Local static and regression verification + +**Files:** + +- Verify all modified Python and test files. + +**Step 1: Run formatting and lint checks** + +```bash +pre-commit run --files \ + python/sglang/srt/layers/cp/kimi_linear.py \ + python/sglang/srt/layers/cp/utils.py \ + python/sglang/srt/models/kimi_linear.py \ + test/registered/cp/test_kimi_linear_cp_v2.py +``` + +**Step 2: Run CP-focused test suites** + +Run the new unit test, existing CP strategy unit test, and any applicable +server-argument tests selected by the diff. + +**Step 3: Inspect the final local diff** + +Require `git diff --check`, review all changes against the design, and confirm +no unrelated user changes are present. + +### Task 6: GB300 end-to-end verification + +**Files:** + +- Remote checkout and logs on `baizhou-dev-2`. +- No generated benchmark artifacts committed to the repository. + +**Step 1: Prepare the devbox** + +Use `rx devbox run baizhou-dev-2`. Clone or update SGLang, fetch the implementation +branch, pull the latest stacked base, and install the current editable Python and +kernel dependencies before launching a job. + +**Step 2: Locate or download the model** + +Use `moonshotai/Kimi-Linear-48B-A3B-Instruct` from a shared cache if present; +otherwise download it to devbox-attached persistent storage. + +**Step 3: Launch the requested configuration** + +Start SGLang with: + +```bash +--tp 4 --attn-cp-size 4 --enable-prefill-cp --cp-strategy zigzag +``` + +Capture the exact command, commit, model path, backend selection, and complete +server log. + +**Step 4: Run GSM8K** + +Use the repository evaluation command against the live endpoint. Record sample +count, accuracy, and any mismatch/error output. If no established Kimi-Linear +threshold exists, compare against a TP4 non-CP run using the same model, +tokenizer, prompts, and decoding settings. + +**Step 5: Diagnose until verified** + +If the server or evaluation fails, preserve the first failure signature, add a +focused regression test where feasible, and repeat local plus GB300 validation. + +### Task 7: Review, commit, push, and open the stacked PR + +**Files:** + +- Review: all files changed since `origin/cp-v2-mla-prefill`. + +**Step 1: Re-run verification before claiming completion** + +Capture fresh output for the focused tests, pre-commit checks, and GSM8K result. + +**Step 2: Commit intentionally** + +Keep the design document commit and create logically scoped implementation/test +commits. Do not fold in changes from the stacked base or unrelated worktree +state. + +**Step 3: Push to the authenticated fork** + +Push `codex/kimi-linear-cp-v2` to `Fridge003/sglang`. + +**Step 4: Open a draft stacked PR** + +Open the PR with base `sgl-project:cp-v2-mla-prefill`. Include the layout +transition table, dependency on #31619, unit-test commands, exact GB300 launch +command, and GSM8K result. From 1d54fc2a454f4decc68ae3fdd7196903ca64ad1d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:36:56 -0700 Subject: [PATCH 03/35] test: specify Kimi-Linear CP-v2 entry gather --- test/registered/cp/test_kimi_linear_cp_v2.py | 55 ++++++++++++++++++++ 1 file changed, 55 insertions(+) create mode 100644 test/registered/cp/test_kimi_linear_cp_v2.py diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py new file mode 100644 index 000000000000..c08d179bf9d8 --- /dev/null +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -0,0 +1,55 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class _RecordingStrategy: + def __init__(self): + self.gather_calls = 0 + + def gather_hidden_states(self, hidden_states, forward_batch, stream=None): + del forward_batch, stream + self.gather_calls += 1 + return hidden_states + 10 + + +class TestKimiLinearCPV2LayerCommunicator(CustomTestCase): + def test_first_kda_layer_gathers_model_entry_shard(self): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=True, + previous_is_kda_layer=None, + ) + hidden_states = torch.arange(4).view(2, 2) + + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, residual = communicator.prepare_attn( + hidden_states, + None, + SimpleNamespace(), + ) + + self.assertEqual(strategy.gather_calls, 1) + torch.testing.assert_close(output, hidden_states + 10) + self.assertIsNone(residual) + + +if __name__ == "__main__": + unittest.main() From 407354fd1a75d44b19728da939f50e682e712c84 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:38:11 -0700 Subject: [PATCH 04/35] feat: gather Kimi-Linear CP input for KDA --- python/sglang/srt/layers/cp/kimi_linear.py | 55 ++++++++++++++++++++++ 1 file changed, 55 insertions(+) create mode 100644 python/sglang/srt/layers/cp/kimi_linear.py diff --git a/python/sglang/srt/layers/cp/kimi_linear.py b/python/sglang/srt/layers/cp/kimi_linear.py new file mode 100644 index 000000000000..fe201c468ea0 --- /dev/null +++ b/python/sglang/srt/layers/cp/kimi_linear.py @@ -0,0 +1,55 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""CP-v2 token-layout transitions for Kimi-Linear decoder layers.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any, Optional, Tuple + +from sglang.srt.layers.cp.utils import get_cp_strategy, is_cp_v2_active + +if TYPE_CHECKING: + import torch + + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +class KimiLinearCPV2LayerCommunicator: + """Convert Kimi-Linear layer inputs between KDA and MLA token layouts.""" + + def __init__( + self, + *, + is_kda_layer: bool, + previous_is_kda_layer: Optional[bool], + ) -> None: + self._gather_before_attn = is_kda_layer and (previous_is_kda_layer is not True) + + def prepare_attn( + self, + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + forward_batch: ForwardBatch, + stream: Optional[Any] = None, + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + if not is_cp_v2_active(forward_batch) or not self._gather_before_attn: + return hidden_states, residual + + strategy = get_cp_strategy() + assert strategy is not None + hidden_states = strategy.gather_hidden_states( + hidden_states, forward_batch, stream + ) + return hidden_states, residual From 3f452cd575c28a5bb536e7e3c0d562486c05d053 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:38:52 -0700 Subject: [PATCH 05/35] test: specify KDA to MLA CP split --- test/registered/cp/test_kimi_linear_cp_v2.py | 35 ++++++++++++++++++++ 1 file changed, 35 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index c08d179bf9d8..83f3e7036e4e 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -14,12 +14,18 @@ class _RecordingStrategy: def __init__(self): self.gather_calls = 0 + self.shard_calls = 0 def gather_hidden_states(self, hidden_states, forward_batch, stream=None): del forward_batch, stream self.gather_calls += 1 return hidden_states + 10 + def shard_hidden_states(self, hidden_states, forward_batch): + del forward_batch + self.shard_calls += 1 + return hidden_states[::2] + class TestKimiLinearCPV2LayerCommunicator(CustomTestCase): def test_first_kda_layer_gathers_model_entry_shard(self): @@ -50,6 +56,35 @@ def test_first_kda_layer_gathers_model_entry_shard(self): torch.testing.assert_close(output, hidden_states + 10) self.assertIsNone(residual) + def test_kda_to_mla_shards_hidden_states_and_residual(self): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=False, + previous_is_kda_layer=True, + ) + hidden_states = torch.arange(8).view(4, 2) + residual = hidden_states + 100 + + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, output_residual = communicator.prepare_attn( + hidden_states, + residual, + SimpleNamespace(), + ) + + self.assertEqual(strategy.shard_calls, 2) + torch.testing.assert_close(output, hidden_states[::2]) + torch.testing.assert_close(output_residual, residual[::2]) + if __name__ == "__main__": unittest.main() From 630f45daa667ca6330a9df10566a85212aa8c20f Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:39:31 -0700 Subject: [PATCH 06/35] feat: shard Kimi-Linear state before MLA --- python/sglang/srt/layers/cp/kimi_linear.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/layers/cp/kimi_linear.py b/python/sglang/srt/layers/cp/kimi_linear.py index fe201c468ea0..3c8d924cd9b8 100644 --- a/python/sglang/srt/layers/cp/kimi_linear.py +++ b/python/sglang/srt/layers/cp/kimi_linear.py @@ -36,6 +36,7 @@ def __init__( previous_is_kda_layer: Optional[bool], ) -> None: self._gather_before_attn = is_kda_layer and (previous_is_kda_layer is not True) + self._shard_before_attn = not is_kda_layer and previous_is_kda_layer is True def prepare_attn( self, @@ -44,12 +45,17 @@ def prepare_attn( forward_batch: ForwardBatch, stream: Optional[Any] = None, ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - if not is_cp_v2_active(forward_batch) or not self._gather_before_attn: + if not is_cp_v2_active(forward_batch): return hidden_states, residual strategy = get_cp_strategy() assert strategy is not None - hidden_states = strategy.gather_hidden_states( - hidden_states, forward_batch, stream - ) + if self._gather_before_attn: + hidden_states = strategy.gather_hidden_states( + hidden_states, forward_batch, stream + ) + elif self._shard_before_attn: + hidden_states = strategy.shard_hidden_states(hidden_states, forward_batch) + if residual is not None: + residual = strategy.shard_hidden_states(residual, forward_batch) return hidden_states, residual From d5921b9e606facbfc64ffe528b948cde1c5376a3 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:40:07 -0700 Subject: [PATCH 07/35] test: specify MLA to KDA CP gather --- test/registered/cp/test_kimi_linear_cp_v2.py | 29 ++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 83f3e7036e4e..91bc7573ba40 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -85,6 +85,35 @@ def test_kda_to_mla_shards_hidden_states_and_residual(self): torch.testing.assert_close(output, hidden_states[::2]) torch.testing.assert_close(output_residual, residual[::2]) + def test_mla_to_kda_gathers_hidden_states_and_residual(self): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=True, + previous_is_kda_layer=False, + ) + hidden_states = torch.arange(4).view(2, 2) + residual = hidden_states + 100 + + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, output_residual = communicator.prepare_attn( + hidden_states, + residual, + SimpleNamespace(), + ) + + self.assertEqual(strategy.gather_calls, 2) + torch.testing.assert_close(output, hidden_states + 10) + torch.testing.assert_close(output_residual, residual + 10) + if __name__ == "__main__": unittest.main() From 3c213d602eddcb6fc566497c1dbea53de2fec5fe Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:40:39 -0700 Subject: [PATCH 08/35] feat: gather Kimi-Linear residual before KDA --- python/sglang/srt/layers/cp/kimi_linear.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/python/sglang/srt/layers/cp/kimi_linear.py b/python/sglang/srt/layers/cp/kimi_linear.py index 3c8d924cd9b8..c74f31de6f0d 100644 --- a/python/sglang/srt/layers/cp/kimi_linear.py +++ b/python/sglang/srt/layers/cp/kimi_linear.py @@ -54,6 +54,10 @@ def prepare_attn( hidden_states = strategy.gather_hidden_states( hidden_states, forward_batch, stream ) + if residual is not None: + residual = strategy.gather_hidden_states( + residual, forward_batch, stream + ) elif self._shard_before_attn: hidden_states = strategy.shard_hidden_states(hidden_states, forward_batch) if residual is not None: From 7b240109fc9a91218a64f5d006b45c88164410b7 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:42:03 -0700 Subject: [PATCH 09/35] test: specify Kimi decoder CP communicator wiring --- test/registered/cp/test_kimi_linear_cp_v2.py | 71 +++++++++++++++++++- 1 file changed, 70 insertions(+), 1 deletion(-) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 91bc7573ba40..ae105b0944cf 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -1,10 +1,11 @@ import unittest from types import SimpleNamespace -from unittest.mock import patch +from unittest.mock import MagicMock, patch import torch from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator +from sglang.srt.models.kimi_linear import KimiDecoderLayer from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -115,5 +116,73 @@ def test_mla_to_kda_gathers_hidden_states_and_residual(self): torch.testing.assert_close(output_residual, residual + 10) +class TestKimiDecoderLayerCPV2Wiring(CustomTestCase): + def test_layer_prepares_cp_layout_before_input_norm(self): + config = SimpleNamespace( + hidden_size=4, + is_moe=False, + intermediate_size=8, + hidden_act="silu", + rms_norm_eps=1e-5, + is_kda_layer=lambda layer_idx: layer_idx != 3, + ) + communicator = MagicMock() + hidden_states = torch.arange(8).view(2, 4) + prepared_hidden_states = hidden_states + 10 + normalized_hidden_states = hidden_states + 20 + attention_output = hidden_states + 30 + post_norm_output = hidden_states + 40 + post_norm_residual = hidden_states + 50 + mlp_output = hidden_states + 60 + communicator.prepare_attn.return_value = (prepared_hidden_states, None) + input_layernorm = MagicMock(return_value=normalized_hidden_states) + self_attn = MagicMock(return_value=attention_output) + post_attention_layernorm = MagicMock( + return_value=(post_norm_output, post_norm_residual) + ) + mlp = MagicMock(return_value=mlp_output) + stream = MagicMock() + forward_batch = SimpleNamespace() + + with ( + patch( + "sglang.srt.models.kimi_linear.KimiLinearCPV2LayerCommunicator", + return_value=communicator, + ) as communicator_cls, + patch( + "sglang.srt.models.kimi_linear.KimiDeltaAttention", + return_value=self_attn, + ), + patch("sglang.srt.models.kimi_linear.KimiMLP", return_value=mlp), + patch( + "sglang.srt.models.kimi_linear.RMSNorm", + side_effect=[input_layernorm, post_attention_layernorm], + ), + patch("torch.cuda.current_stream", return_value=stream), + ): + layer = KimiDecoderLayer(config=config, layer_idx=3) + output, output_residual = layer( + positions=torch.arange(2), + hidden_states=hidden_states, + forward_batch=forward_batch, + residual=None, + zero_allocator=MagicMock(), + ) + + communicator_cls.assert_called_once_with( + is_kda_layer=False, + previous_is_kda_layer=True, + ) + communicator.prepare_attn.assert_called_once_with( + hidden_states, + None, + forward_batch, + stream, + ) + input_layernorm.assert_called_once_with(prepared_hidden_states) + torch.testing.assert_close(output, mlp_output) + torch.testing.assert_close(output_residual, post_norm_residual) + + if __name__ == "__main__": unittest.main() From 2d7e74dbf50d074fb0603a52b15e808d77069726 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:43:13 -0700 Subject: [PATCH 10/35] feat: wire CP-v2 transitions into Kimi layers --- python/sglang/srt/models/kimi_linear.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 9d51b3f77bb3..758ae95f9450 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -16,6 +16,7 @@ tensor_model_parallel_all_reduce, ) from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder +from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelBatchedLinear, @@ -423,8 +424,15 @@ def __init__( self.alt_stream = alt_stream self.is_moe = config.is_moe + is_kda_layer = config.is_kda_layer(layer_idx) + self.cp_communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=is_kda_layer, + previous_is_kda_layer=( + config.is_kda_layer(layer_idx - 1) if layer_idx > 0 else None + ), + ) - if config.is_kda_layer(layer_idx): + if is_kda_layer: self.self_attn = KimiDeltaAttention( layer_idx=layer_idx, hidden_size=config.hidden_size, @@ -483,6 +491,13 @@ def forward( residual: Optional[torch.Tensor], zero_allocator: BumpAllocator, ) -> tuple[torch.Tensor, torch.Tensor]: + hidden_states, residual = self.cp_communicator.prepare_attn( + hidden_states, + residual, + forward_batch, + torch.cuda.current_stream(), + ) + # Self Attention if residual is None: residual = hidden_states From 7f9362e4bc282b84d75063bb065588bbe9df2b86 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:44:00 -0700 Subject: [PATCH 11/35] test: complete Kimi decoder fixture --- test/registered/cp/test_kimi_linear_cp_v2.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index ae105b0944cf..bccd53ba14df 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -120,6 +120,12 @@ class TestKimiDecoderLayerCPV2Wiring(CustomTestCase): def test_layer_prepares_cp_layout_before_input_norm(self): config = SimpleNamespace( hidden_size=4, + num_attention_heads=4, + qk_nope_head_dim=2, + qk_rope_head_dim=2, + v_head_dim=2, + q_lora_rank=None, + kv_lora_rank=2, is_moe=False, intermediate_size=8, hidden_act="silu", From 92ce1afce18cb052f8631019679a0b8413c429c0 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:44:35 -0700 Subject: [PATCH 12/35] test: isolate Kimi MLA layer wiring --- test/registered/cp/test_kimi_linear_cp_v2.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index bccd53ba14df..689c088d7966 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -159,6 +159,10 @@ def test_layer_prepares_cp_layout_before_input_norm(self): "sglang.srt.models.kimi_linear.KimiDeltaAttention", return_value=self_attn, ), + patch( + "sglang.srt.models.kimi_linear.KimiMLAAttention", + return_value=self_attn, + ), patch("sglang.srt.models.kimi_linear.KimiMLP", return_value=mlp), patch( "sglang.srt.models.kimi_linear.RMSNorm", From 70a876789d9a0012beb1e8e6d667cf85f2b86eea Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:45:18 -0700 Subject: [PATCH 13/35] test: require CP-v2 for Kimi-Linear --- test/registered/cp/test_kimi_linear_cp_v2.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 689c088d7966..2e456f2fe8fb 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -5,6 +5,7 @@ import torch from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator +from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES from sglang.srt.models.kimi_linear import KimiDecoderLayer from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -194,5 +195,10 @@ def test_layer_prepares_cp_layout_before_input_norm(self): torch.testing.assert_close(output_residual, post_norm_residual) +class TestKimiLinearCPV2Activation(CustomTestCase): + def test_kimi_linear_uses_cp_v2_by_default(self): + self.assertIn("KimiLinearForCausalLM", CP_V2_DEFAULT_MODEL_CLASSES) + + if __name__ == "__main__": unittest.main() From 898ac515a2c345d398cf81c8610245350caff68d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:45:47 -0700 Subject: [PATCH 14/35] feat: enable CP-v2 for Kimi-Linear --- python/sglang/srt/layers/cp/utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python/sglang/srt/layers/cp/utils.py b/python/sglang/srt/layers/cp/utils.py index 1ba7545cac81..f9bc3f2b2877 100644 --- a/python/sglang/srt/layers/cp/utils.py +++ b/python/sglang/srt/layers/cp/utils.py @@ -41,6 +41,7 @@ { "Qwen3MoeForCausalLM", "DeepseekV3ForCausalLM", + "KimiLinearForCausalLM", } ) From cd717800cfb3bf8098bf85112abe1d28a7f17aee Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:46:26 -0700 Subject: [PATCH 15/35] test: require Kimi input embedding accessor --- test/registered/cp/test_kimi_linear_cp_v2.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 2e456f2fe8fb..32f231d879a2 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -6,7 +6,7 @@ from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES -from sglang.srt.models.kimi_linear import KimiDecoderLayer +from sglang.srt.models.kimi_linear import KimiDecoderLayer, KimiLinearForCausalLM from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -199,6 +199,17 @@ class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): self.assertIn("KimiLinearForCausalLM", CP_V2_DEFAULT_MODEL_CLASSES) + def test_causal_lm_exposes_input_embeddings(self): + causal_lm = object.__new__(KimiLinearForCausalLM) + embeddings = MagicMock() + object.__setattr__( + causal_lm, + "model", + SimpleNamespace(embed_tokens=embeddings), + ) + + self.assertIs(causal_lm.get_input_embeddings(), embeddings) + if __name__ == "__main__": unittest.main() From 8754d7ebe50e60761cb8874e9c3860b7b33991fb Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:46:57 -0700 Subject: [PATCH 16/35] feat: expose Kimi input embeddings to CP-v2 --- python/sglang/srt/models/kimi_linear.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 758ae95f9450..a0b3bcc62614 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -654,6 +654,9 @@ def __init__( logit_scale = getattr(self.config, "logit_scale", 1.0) self.logits_processor = LogitsProcessor(config=config, logit_scale=logit_scale) + def get_input_embeddings(self) -> nn.Module: + return self.model.embed_tokens + @torch.no_grad() def forward( self, From 78c208374f2ebbe95bcea3ff743f51a902342cd3 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 21:49:24 -0700 Subject: [PATCH 17/35] test: cover Kimi CP-v2 no-op and zigzag round trip --- test/registered/cp/test_kimi_linear_cp_v2.py | 141 +++++++++++++++++++ 1 file changed, 141 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 32f231d879a2..2991ac1e6d7d 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -6,7 +6,9 @@ from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES +from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy from sglang.srt.models.kimi_linear import KimiDecoderLayer, KimiLinearForCausalLM +from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -29,6 +31,17 @@ def shard_hidden_states(self, hidden_states, forward_batch): return hidden_states[::2] +class _SequencedFakeCPGroup: + def __init__(self, *rank_tensor_sets): + self.rank_tensor_sets = rank_tensor_sets + self.call_index = 0 + + def cp_all_gather_into_tensor_async(self, output, input_tensor, stream): + del input_tensor, stream + torch.cat(self.rank_tensor_sets[self.call_index], dim=0, out=output) + self.call_index += 1 + + class TestKimiLinearCPV2LayerCommunicator(CustomTestCase): def test_first_kda_layer_gathers_model_entry_shard(self): strategy = _RecordingStrategy() @@ -116,6 +129,134 @@ def test_mla_to_kda_gathers_hidden_states_and_residual(self): torch.testing.assert_close(output, hidden_states + 10) torch.testing.assert_close(output_residual, residual + 10) + def test_same_layout_and_inactive_cp_v2_are_noops(self): + hidden_states = torch.arange(8).view(4, 2) + residual = hidden_states + 100 + + for is_kda_layer, previous_is_kda_layer, cp_v2_active in ( + (True, True, True), + (False, False, True), + (False, None, True), + (True, False, False), + (False, True, False), + ): + with self.subTest( + is_kda_layer=is_kda_layer, + previous_is_kda_layer=previous_is_kda_layer, + cp_v2_active=cp_v2_active, + ): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=is_kda_layer, + previous_is_kda_layer=previous_is_kda_layer, + ) + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=cp_v2_active, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, output_residual = communicator.prepare_attn( + hidden_states, + residual, + SimpleNamespace(), + ) + + self.assertIs(output, hidden_states) + self.assertIs(output_residual, residual) + self.assertEqual(strategy.gather_calls, 0) + self.assertEqual(strategy.shard_calls, 0) + + def test_zigzag_kda_mla_kda_round_trip_restores_token_order(self): + cp_size = 4 + seq_lens = [11, 13] + extend_seq_lens = [9, 10] + num_tokens = sum(extend_seq_lens) + hidden_states = torch.arange(num_tokens * 2).view(num_tokens, 2) + residual = hidden_states + 100 + strategy = ZigzagCPStrategy(cp_size=cp_size) + shard_communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=False, + previous_is_kda_layer=True, + ) + gather_communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=True, + previous_is_kda_layer=False, + ) + + forward_batches = [] + local_hidden_states = [] + local_residuals = [] + for rank in range(cp_size): + with get_parallel().override(attn_cp_rank=rank): + metadata = strategy.build_metadata( + num_tokens=num_tokens, + seqs_len=seq_lens, + extend_seqs_len=extend_seq_lens, + ) + forward_batch = SimpleNamespace(attn_cp_metadata=metadata) + forward_batches.append(forward_batch) + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + local_hidden, local_residual = shard_communicator.prepare_attn( + hidden_states, + residual, + forward_batch, + ) + local_hidden_states.append(local_hidden) + local_residuals.append(local_residual) + + max_rank_len = forward_batches[0].attn_cp_metadata.max_rank_len[0] + + def _pad_rank_tensors(rank_tensors): + return [ + torch.nn.functional.pad( + tensor, + [0, 0, 0, max_rank_len - tensor.shape[0]], + ) + for tensor in rank_tensors + ] + + padded_hidden_states = _pad_rank_tensors(local_hidden_states) + padded_residuals = _pad_rank_tensors(local_residuals) + + for rank in range(cp_size): + group = _SequencedFakeCPGroup( + padded_hidden_states, + padded_residuals, + ) + with ( + get_parallel().override(attn_cp_group=group), + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + gathered_hidden, gathered_residual = gather_communicator.prepare_attn( + local_hidden_states[rank], + local_residuals[rank], + forward_batches[rank], + ) + + torch.testing.assert_close(gathered_hidden, hidden_states) + torch.testing.assert_close(gathered_residual, residual) + class TestKimiDecoderLayerCPV2Wiring(CustomTestCase): def test_layer_prepares_cp_layout_before_input_norm(self): From 02f464bae2c2709b5931ae4ce9476c9319ffe8d8 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:00:10 -0700 Subject: [PATCH 18/35] test: require KDA heads to use global TP --- test/registered/cp/test_kimi_linear_cp_v2.py | 62 +++++++++++++++++++- 1 file changed, 61 insertions(+), 1 deletion(-) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 2991ac1e6d7d..ebffcd16b834 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -7,7 +7,11 @@ from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy -from sglang.srt.models.kimi_linear import KimiDecoderLayer, KimiLinearForCausalLM +from sglang.srt.models.kimi_linear import ( + KimiDecoderLayer, + KimiDeltaAttention, + KimiLinearForCausalLM, +) from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase @@ -335,6 +339,62 @@ def test_layer_prepares_cp_layout_before_input_norm(self): torch.testing.assert_close(output, mlp_output) torch.testing.assert_close(output_residual, post_norm_residual) + def test_kda_backend_partitions_heads_over_global_tp(self): + config = SimpleNamespace( + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + }, + v_head_dim=128, + dtype=torch.bfloat16, + ) + radix_linear_attention = MagicMock() + + with ( + get_parallel().override( + tp_size=4, + tp_rank=2, + attn_tp_size=1, + attn_tp_rank=0, + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelRepeatedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.ColumnParallelBatchedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.FusedRMSNormGated", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RowParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RadixLinearAttention", + return_value=radix_linear_attention, + ) as radix_linear_attention_cls, + ): + KimiDeltaAttention( + layer_idx=0, + hidden_size=256, + config=config, + ) + + radix_linear_attention_cls.assert_called_once() + backend_args = radix_linear_attention_cls.call_args.kwargs + self.assertEqual(backend_args["num_q_heads"], 8) + self.assertEqual(backend_args["num_k_heads"], 8) + self.assertEqual(backend_args["num_v_heads"], 8) + class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): From 4f6220ae32bcfe21363237a3505759ea1773a543 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:00:46 -0700 Subject: [PATCH 19/35] fix: partition KDA backend heads over global TP --- python/sglang/srt/models/kimi_linear.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index a0b3bcc62614..64c0f9807727 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -318,9 +318,9 @@ def __init__( self.attn = RadixLinearAttention( layer_id=self.layer_idx, - num_q_heads=self.num_k_heads // self.attn_tp_size, - num_k_heads=self.num_k_heads // self.attn_tp_size, - num_v_heads=self.num_v_heads // self.attn_tp_size, + num_q_heads=self.num_k_heads // self.tp_size, + num_k_heads=self.num_k_heads // self.tp_size, + num_v_heads=self.num_v_heads // self.tp_size, head_q_dim=self.head_k_dim, head_k_dim=self.head_k_dim, head_v_dim=self.head_v_dim, From bf733d958a6f5c100cccac26df7d0d760978f3a1 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:01:33 -0700 Subject: [PATCH 20/35] test: require unfused KDA projection to use global TP --- test/registered/cp/test_kimi_linear_cp_v2.py | 60 ++++++++++++++++++++ 1 file changed, 60 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index ebffcd16b834..cfd04bc62d2b 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -395,6 +395,66 @@ def test_kda_backend_partitions_heads_over_global_tp(self): self.assertEqual(backend_args["num_k_heads"], 8) self.assertEqual(backend_args["num_v_heads"], 8) + def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): + config = SimpleNamespace( + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + }, + v_head_dim=128, + dtype=torch.bfloat16, + ) + qkv_parallel_linear = MagicMock() + + with ( + get_parallel().override( + tp_size=4, + tp_rank=2, + attn_tp_size=1, + attn_tp_rank=0, + ), + patch( + "sglang.srt.models.kimi_linear.QKVParallelLinear", + return_value=qkv_parallel_linear, + ) as qkv_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.ReplicatedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.ColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.FusedRMSNormGated", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RowParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RadixLinearAttention", + return_value=MagicMock(), + ), + ): + KimiDeltaAttention( + layer_idx=0, + hidden_size=256, + config=config, + quant_config=MagicMock(), + ) + + qkv_parallel_linear_cls.assert_called_once() + projection_args = qkv_parallel_linear_cls.call_args.kwargs + self.assertEqual(projection_args["tp_rank"], 2) + self.assertEqual(projection_args["tp_size"], 4) + class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): From f0364b563066f881b20b5bb58f93b6e4aff00bbe Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:02:24 -0700 Subject: [PATCH 21/35] fix: shard unfused KDA projections over global TP --- python/sglang/srt/models/kimi_linear.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 64c0f9807727..f1fcbc621d99 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -178,7 +178,6 @@ def __init__( ) -> None: super().__init__() self.tp_size = get_parallel().tp_size - self.attn_tp_size = get_parallel().attn_tp_size self.hidden_size = hidden_size self.config = config self.head_dim = config.linear_attn_config["head_dim"] @@ -225,7 +224,7 @@ def __init__( ) else: # Unfused path: separate QKVParallelLinear - attn_tp_rank = get_parallel().attn_tp_rank + tp_rank = get_parallel().tp_rank self.qkv_proj = QKVParallelLinear( self.hidden_size, self.head_dim, @@ -233,8 +232,8 @@ def __init__( self.num_k_heads, bias=False, quant_config=quant_config, - tp_rank=attn_tp_rank, - tp_size=self.attn_tp_size, + tp_rank=tp_rank, + tp_size=self.tp_size, v_head_size=self.head_v_dim, prefix=f"{prefix}.qkv_proj", ) From 01b92d04a5145fce14f8214ac93b7fc577e7d5be Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:04:46 -0700 Subject: [PATCH 22/35] test: require KDA cache state to use global TP --- test/registered/cp/test_kimi_linear_cp_v2.py | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index cfd04bc62d2b..831a4701de8a 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -4,6 +4,7 @@ import torch +from sglang.srt.configs.kimi_linear import KimiLinearConfig from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -455,6 +456,25 @@ def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): self.assertEqual(projection_args["tp_rank"], 2) self.assertEqual(projection_args["tp_size"], 4) + def test_kda_cache_shape_uses_global_tp_size(self): + config = KimiLinearConfig( + num_hidden_layers=2, + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + "kda_layers": [1], + "full_attn_layers": [2], + }, + ) + + with get_parallel().override(tp_size=4, attn_tp_size=1): + shape = config.mamba2_cache_params.shape + + self.assertEqual(shape.temporal, (8, 128, 128)) + self.assertEqual(shape.conv, [(3, 3072)]) + self.assertEqual(shape.num_k_heads_per_tp, 8) + class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): From e2ecec60efc207343df2d8ed1fb2da02d9efed35 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:05:31 -0700 Subject: [PATCH 23/35] fix: shard KDA state cache over global TP --- python/sglang/srt/configs/kimi_linear.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/sglang/srt/configs/kimi_linear.py b/python/sglang/srt/configs/kimi_linear.py index 181dca1d3ec1..a26c947bd0d7 100644 --- a/python/sglang/srt/configs/kimi_linear.py +++ b/python/sglang/srt/configs/kimi_linear.py @@ -154,7 +154,7 @@ def full_attention_layer_ids(self): def mamba2_cache_params(self) -> KimiLinearCacheParams: shape = KimiLinearStateShape.create( - tp_world_size=get_parallel().attn_tp_size, + tp_world_size=get_parallel().tp_size, num_heads=self.linear_attn_config["num_heads"], head_dim=self.linear_attn_config["head_dim"], conv_kernel_size=self.linear_attn_config["short_conv_kernel_size"], From 55f2a644cd70157d9af960a9b0a08430f1a83f77 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:09:53 -0700 Subject: [PATCH 24/35] test: cover Kimi CP-v2 embedding keyword --- test/registered/cp/test_kimi_linear_cp_v2.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 831a4701de8a..cd073b88007c 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -1,3 +1,4 @@ +import inspect import unittest from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -12,6 +13,7 @@ KimiDecoderLayer, KimiDeltaAttention, KimiLinearForCausalLM, + KimiLinearModel, ) from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci @@ -491,6 +493,11 @@ def test_causal_lm_exposes_input_embeddings(self): self.assertIs(causal_lm.get_input_embeddings(), embeddings) + def test_inner_model_accepts_cp_v2_input_embeds_keyword(self): + parameters = inspect.signature(KimiLinearModel.forward).parameters + + self.assertIn("input_embeds", parameters) + if __name__ == "__main__": unittest.main() From b41d63a7f9f12d608a9a99c6a1c47876b74cc773 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:10:42 -0700 Subject: [PATCH 25/35] fix: accept CP-v2 input embeddings in Kimi --- python/sglang/srt/models/kimi_linear.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index f1fcbc621d99..b3ea15f7aee4 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -572,12 +572,12 @@ def forward( input_ids: torch.Tensor | None, positions: torch.Tensor, forward_batch: ForwardBatch, - inputs_embeds: torch.Tensor | None = None, + input_embeds: torch.Tensor | None = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: if get_pp_group().is_first_rank: - if inputs_embeds is not None: - hidden_states = inputs_embeds + if input_embeds is not None: + hidden_states = input_embeds else: hidden_states = self.embed_tokens(input_ids) residual = None @@ -662,14 +662,14 @@ def forward( input_ids: torch.Tensor, positions: torch.Tensor, forward_batch: ForwardBatch, - inputs_embeds: Optional[torch.Tensor] = None, + input_embeds: Optional[torch.Tensor] = None, pp_proxy_tensors: Optional[PPProxyTensors] = None, ) -> torch.Tensor: hidden_states = self.model( input_ids, positions, forward_batch, - inputs_embeds, + input_embeds, pp_proxy_tensors, ) if self.pp_group.is_last_rank: From 5af9dd862c060001c26c4ce37d568603e7ef2ca5 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:19:29 -0700 Subject: [PATCH 26/35] test: cover FlashInfer MLA CP-v2 dispatch --- test/registered/cp/test_kimi_linear_cp_v2.py | 56 ++++++++++++++++++++ 1 file changed, 56 insertions(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index cd073b88007c..2828771af645 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -6,6 +6,9 @@ import torch from sglang.srt.configs.kimi_linear import KimiLinearConfig +from sglang.srt.layers.attention.flashinfer_mla_backend import ( + FlashInferMLAAttnBackend, +) from sglang.srt.layers.cp.kimi_linear import KimiLinearCPV2LayerCommunicator from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES from sglang.srt.layers.cp.zigzag import ZigzagCPStrategy @@ -499,5 +502,58 @@ def test_inner_model_accepts_cp_v2_input_embeds_keyword(self): self.assertIn("input_embeds", parameters) +class TestKimiLinearFlashInferMLACP(CustomTestCase): + def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): + backend = object.__new__(FlashInferMLAAttnBackend) + backend.device = torch.device("cpu") + backend._get_cp_prefill_metadata = MagicMock( + return_value=SimpleNamespace(wrappers=[MagicMock(), MagicMock()]) + ) + backend._run_cp_paged_attention = MagicMock() + strategy = MagicMock() + expected = torch.randn(5, 4) + strategy.run_attention.return_value = expected + q = torch.randn(5, 8) + q_rope = torch.randn(5, 4) + k = torch.randn(5, 8) + v = torch.randn(5, 8) + k_rope = torch.randn(5, 4) + layer = SimpleNamespace( + tp_q_head_num=2, + v_head_dim=4, + head_dim=6, + ) + forward_batch = SimpleNamespace() + + with ( + patch( + "sglang.srt.layers.attention.flashinfer_mla_backend.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.attention.flashinfer_mla_backend.get_cp_strategy", + return_value=strategy, + ), + ): + output = backend.forward_extend( + q, + k, + v, + layer, + forward_batch, + q_rope=q_rope, + k_rope=k_rope, + ) + + strategy.materialize_full_mla_kv.assert_called_once_with( + forward_batch, + layer, + k, + k_rope, + ) + strategy.run_attention.assert_called_once() + self.assertIs(output, expected) + + if __name__ == "__main__": unittest.main() From 34ced26ffdba31b67be380a63d79fef62e3e7034 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:21:03 -0700 Subject: [PATCH 27/35] feat: support FlashInfer MLA in CP-v2 --- .../attention/flashinfer_mla_backend.py | 159 ++++++++++++++++++ 1 file changed, 159 insertions(+) diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 0d68b303c7e9..853151ef8547 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -23,6 +23,8 @@ from sglang.srt.layers.attention.flashinfer_backend import ( create_flashinfer_kv_indices_triton, ) +from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy +from sglang.srt.layers.cp.utils import is_cp_v2_active from sglang.srt.layers.dcp import ( DecodeContextParallelMetadata, update_local_kv_lens_for_dcp, @@ -77,6 +79,13 @@ class PrefillMetadata: use_ragged: bool +@dataclass +class CPPrefillMetadata: + wrappers: tuple[BatchMLAPagedAttentionWrapper, BatchMLAPagedAttentionWrapper] + kv_indptrs: tuple[torch.Tensor, torch.Tensor] + kv_indices: tuple[torch.Tensor, torch.Tensor] + + # Reuse this workspace buffer across all flashinfer wrappers @@ -299,6 +308,7 @@ def __init__( # Other metadata self.forward_metadata: Union[PrefillMetadata, DecodeMetadata] = None + self.cp_prefill_metadata: Optional[CPPrefillMetadata] = None self.decode_cuda_graph_metadata = {} self.prefill_cuda_graph_metadata = {} # For verify @@ -377,6 +387,7 @@ def init_forward_metadata_out_graph( ) def init_forward_metadata(self, forward_batch: ForwardBatch): + self.cp_prefill_metadata = None if forward_batch.forward_mode.is_decode_or_idle(): self.indices_updater_decode.update( forward_batch.req_pool_indices, @@ -511,6 +522,142 @@ def init_mha_chunk_metadata( """Init the metadata for a forward pass.""" self.mha_chunk_kv_cache.update_wrapper(forward_batch, disable_flashinfer_ragged) + def _plan_cp_prefill_wrapper( + self, + forward_batch: ForwardBatch, + qo_indptr: torch.Tensor, + kv_lens: torch.Tensor, + kv_lens_sum: int, + ): + bs = len(forward_batch.req_pool_indices) + kv_indptr = torch.zeros( + bs + 1, + dtype=torch.int32, + device=forward_batch.req_pool_indices.device, + ) + kv_indptr[1:] = torch.cumsum(kv_lens, dim=0) + kv_indices = torch.empty( + kv_lens_sum, + dtype=torch.int32, + device=forward_batch.req_pool_indices.device, + ) + req_to_token = self.req_to_token_pool.req_to_token + create_flashinfer_kv_indices_triton[(bs,)]( + req_to_token, + forward_batch.req_pool_indices, + kv_lens, + kv_indptr, + None, + kv_indices, + req_to_token.shape[1], + ) + + wrapper = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, + backend="auto", + ) + updater = self.indices_updater_prefill + wrapper.plan( + qo_indptr, + kv_indptr, + kv_indices, + kv_lens, + updater.num_local_heads, + updater.kv_lora_rank, + updater.qk_rope_head_dim, + self.page_size, + True, + updater.scaling, + updater.q_data_type, + updater.data_type, + ) + return wrapper, kv_indptr, kv_indices + + def _get_cp_prefill_metadata( + self, forward_batch: ForwardBatch + ) -> CPPrefillMetadata: + if self.cp_prefill_metadata is not None: + return self.cp_prefill_metadata + + meta = forward_batch.attn_cp_metadata + prev = self._plan_cp_prefill_wrapper( + forward_batch, + meta.cu_seqlens_q_prev_tensor, + meta.kv_len_prev_tensor, + sum(meta.kv_len_prev_list), + ) + next_ = self._plan_cp_prefill_wrapper( + forward_batch, + meta.cu_seqlens_q_next_tensor, + meta.kv_len_next_tensor, + sum(meta.kv_len_next_list), + ) + self.cp_prefill_metadata = CPPrefillMetadata( + wrappers=(prev[0], next_[0]), + kv_indptrs=(prev[1], next_[1]), + kv_indices=(prev[2], next_[2]), + ) + return self.cp_prefill_metadata + + def _run_cp_paged_attention( + self, + wrapper: BatchMLAPagedAttentionWrapper, + q: torch.Tensor, + layer: RadixAttention, + ) -> torch.Tensor: + q_nope = q[..., : layer.v_head_dim] + q_rope = q[..., layer.v_head_dim :] + kv_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id).to(q.dtype) + ckv_cache = kv_buffer[:, :, : layer.v_head_dim] + kpe_cache = kv_buffer[:, :, layer.v_head_dim :] + output = q_nope.new_empty(q_nope.shape) + return wrapper.run(q_nope, q_rope, ckv_cache, kpe_cache, out=output) + + def _forward_extend_cp_v2( + self, + q: torch.Tensor, + k: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + save_kv_cache: bool, + q_rope: Optional[torch.Tensor], + k_rope: Optional[torch.Tensor], + ) -> torch.Tensor: + strategy = get_cp_strategy() + assert strategy is not None + assert k_rope is not None + if save_kv_cache: + strategy.materialize_full_mla_kv(forward_batch, layer, k, k_rope) + + if q_rope is None: + q_fused = q.view(-1, layer.tp_q_head_num, layer.head_dim) + else: + q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim) + q_rope = q_rope.view( + -1, + layer.tp_q_head_num, + layer.head_dim - layer.v_head_dim, + ) + q_fused = torch.cat([q_nope, q_rope], dim=-1) + + cp_metadata = self._get_cp_prefill_metadata(forward_batch) + wrapper_index = 0 + + def _mla_cp_attn(q_chunk, *_): + nonlocal wrapper_index + wrapper = cp_metadata.wrappers[wrapper_index] + wrapper_index += 1 + return self._run_cp_paged_attention(wrapper, q_chunk, layer) + + output = strategy.run_attention( + q_fused, + forward_batch, + self.device, + _mla_cp_attn, + attention_backend=CPAttentionBackendKind.FLASH_ATTENTION, + ) + return output.view(-1, layer.tp_q_head_num * layer.v_head_dim) + def forward_extend( self, q: torch.Tensor, @@ -522,6 +669,18 @@ def forward_extend( q_rope: Optional[torch.Tensor] = None, k_rope: Optional[torch.Tensor] = None, ): + if is_cp_v2_active(forward_batch): + assert k is not None and v is not None + return self._forward_extend_cp_v2( + q, + k, + layer, + forward_batch, + save_kv_cache, + q_rope, + k_rope, + ) + if forward_batch.attn_attend_prefix_cache is not None and any( forward_batch.extend_prefix_lens_cpu ): # MHA Chunk From 5ff97cc1feca5a4e10c8639e12a32774c9047463 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Fri, 17 Jul 2026 22:21:40 -0700 Subject: [PATCH 28/35] test: match FlashInfer MLA output layout --- test/registered/cp/test_kimi_linear_cp_v2.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 2828771af645..c8c19d1ec504 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -511,7 +511,7 @@ def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): ) backend._run_cp_paged_attention = MagicMock() strategy = MagicMock() - expected = torch.randn(5, 4) + expected = torch.randn(5, 2, 4) strategy.run_attention.return_value = expected q = torch.randn(5, 8) q_rope = torch.randn(5, 4) @@ -552,7 +552,7 @@ def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): k_rope, ) strategy.run_attention.assert_called_once() - self.assertIs(output, expected) + torch.testing.assert_close(output, expected.view(5, 8)) if __name__ == "__main__": From 37c791fd37b3da3e353f26ff31081aff9c3b3f6c Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 18 Jul 2026 00:02:23 -0700 Subject: [PATCH 29/35] fix: run Kimi KDA and MLP on global TP batches --- .../plans/2026-07-17-kimi-linear-cp-v2.md | 29 +++-- .../2026-07-17-kimi-linear-cp-v2-design.md | 51 ++++++-- python/sglang/srt/layers/cp/kimi_linear.py | 47 ++++++- python/sglang/srt/models/kimi_linear.py | 40 +++++- test/registered/cp/test_kimi_linear_cp_v2.py | 120 ++++++++++++++++-- 5 files changed, 250 insertions(+), 37 deletions(-) diff --git a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md index 171c3094b910..12f17ae68d66 100644 --- a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md +++ b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md @@ -9,9 +9,9 @@ on complete token batches replicated across the four TP ranks. **Architecture:** A Kimi-specific CP-v2 layer communicator converts `hidden_states` and `residual` at KDA/MLA boundaries using the active -`ContextParallelStrategy`. `KimiDecoderLayer` invokes it before input RMSNorm. -The existing CP-v2 eager runner continues to own model-entry splitting and -model-exit gathering. +`ContextParallelStrategy`. `KimiDecoderLayer` invokes it before input RMSNorm +and after MLA attention, before the TP MLP. The existing CP-v2 eager runner +continues to own model-entry splitting and model-exit gathering. **Tech Stack:** Python, PyTorch distributed collectives, SGLang CP-v2 strategy API, `unittest`, four NVIDIA GB300 GPUs, GSM8K evaluation. @@ -77,8 +77,8 @@ Run the same unit-test command and require it to pass. Add and run tests for: - KDA to MLA: shard `hidden_states` and `residual`. -- MLA to KDA: gather `hidden_states` and `residual`. -- KDA to KDA and MLA to MLA: identity/no strategy calls. +- MLA attention to MLP: gather `hidden_states` and `residual`. +- KDA to KDA: identity/no strategy calls. - CP-v2 inactive: identity/no strategy calls. For each case, first observe failure, then implement the smallest transition @@ -87,8 +87,8 @@ logic needed to pass it. **Step 4: Add a real zigzag round-trip test** Use `ZigzagCPStrategy` metadata and the existing fake CP group pattern to prove -that a full tensor split for MLA and gathered for KDA returns to original token -order, including its residual tensor. +that a full tensor split for MLA and gathered before its MLP returns to original +token order, including its residual tensor. **Step 5: Run communicator and existing strategy tests** @@ -120,6 +120,10 @@ In `KimiDecoderLayer.__init__`, compute current and previous layer types from `KimiLinearConfig.is_kda_layer`. Construct the communicator. At the very start of `forward`, call `prepare_attn` before input RMSNorm. +After MLA attention, call `prepare_mlp` before post-attention RMSNorm and the +TP MLP. Shard the final full layer output so the generic model-exit gather keeps +its CP-v2 contract. + **Step 3: Run the focused test and confirm GREEN** ```bash @@ -154,6 +158,15 @@ python -m unittest \ Expected: all pass. +### Task 4b: Keep KDA on global tensor parallelism + +Use global TP rank/size for KDA projections, recurrent cache shape, head +partitioning, and `A_log`/`dt_bias` weight loading. Add a regression test where +`tp_rank` differs from `attn_tp_rank`. + +Add the FlashInfer MLA CP-v2 path because it is the GB300 default backend, and +test its latent-KV materialization plus zigzag dispatch. + ### Task 5: Local static and regression verification **Files:** @@ -241,7 +254,7 @@ state. Push `codex/kimi-linear-cp-v2` to `Fridge003/sglang`. -**Step 4: Open a draft stacked PR** +**Step 4: Open a stacked PR** Open the PR with base `sgl-project:cp-v2-mla-prefill`. Include the layout transition table, dependency on #31619, unit-test commands, exact GB300 launch diff --git a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md index 830ac41bfde7..6a2150ec2976 100644 --- a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md +++ b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md @@ -21,7 +21,7 @@ layers. They must always use the same token layout. - Support Kimi-Linear prefill with `--tp 4 --attn-cp-size 4 --enable-prefill-cp --cp-strategy zigzag`. -- Gather CP-sharded layer state before each KDA region. +- Gather CP-sharded MLA attention output before its post-attention norm and MLP. - Split replicated layer state before each MLA region. - Reuse the active CP strategy for ordering, padding, and collectives. - Preserve the non-CP and decode paths. @@ -29,7 +29,7 @@ layers. They must always use the same token layout. ## Non-goals -- Adding a new CP strategy or attention backend. +- Adding a new CP strategy. - Enabling the unfinished interleave strategy. - Changing KDA kernels, MLA kernels, or KV-cache semantics. - Supporting Kimi K2.5's multimodal wrapper in this change. @@ -48,27 +48,33 @@ Each `KimiDecoderLayer` constructs the communicator with: - Whether the current layer is KDA. - Whether the preceding global layer is KDA. Layer zero has no preceding layer. +- Whether it is the final decoder layer. At the beginning of `KimiDecoderLayer.forward`, the communicator receives `hidden_states`, `residual`, and `forward_batch`, and returns state in the layout -required by the current layer. +required by the current attention layer. After attention, MLA output and its +residual are gathered before the post-attention norm and MLP. Consequently, +every MLP runs on the complete token batch and every layer exits its MLP in the +replicated TP layout. The final layer shards that output once more so the +generic CP-v2 model-exit gather retains its normal contract. It is active only when `is_cp_v2_active(forward_batch)` is true. Otherwise it returns its inputs without communication. ### Transition table -| Incoming state | Current layer | Operation | +| Boundary | Incoming state | Operation | | --- | --- | --- | -| Model-entry CP shard | First KDA | Gather to complete token order | -| Replicated KDA output | KDA | No-op | -| Replicated KDA output | MLA | Split with the active CP strategy | -| CP-sharded MLA output | MLA | No-op | -| CP-sharded MLA output | KDA | Gather to complete token order | +| Model entry → first KDA | CP shard | Gather to complete token order | +| MLP → KDA | Replicated | No-op | +| MLP → MLA | Replicated | Split with the active CP strategy | +| MLA attention → MLP | CP shard | Gather to complete token order | +| Final MLP → model exit | Replicated | Split for the existing model-exit gather | The first Kimi-Linear layer is KDA, so model-entry embeddings are gathered -before layer zero. The final Kimi-Linear layer is MLA, so it remains CP-sharded; -the CP-v2 eager runner performs the existing model-exit gather before logits. +before layer zero. Gathering immediately after each MLA attention operation is +required because the following MoE/MLP contains TP all-reduces whose token +dimensions must be identical on all ranks. ### Hidden state and residual @@ -80,6 +86,23 @@ split uses `ContextParallelStrategy.shard_hidden_states`. The initial layer has `residual=None`; only `hidden_states` is gathered there. No residual addition is moved across the transition. +### KDA tensor parallel state + +KDA remains tensor parallel over the global TP group while MLA attention uses +the CP-sharded layout. Its projections, recurrent-cache head shape, and +`A_log`/`dt_bias` state-parameter loading therefore use global `tp_size` and +`tp_rank`, not attention TP. Under `--tp 4 --attn-cp-size 4`, attention TP has +size one; using its rank would incorrectly load rank zero's recurrent +parameters on all four KDA ranks. + +### FlashInfer MLA + +GB300 selects FlashInfer as the default MLA backend. CP-v2 plans one paged MLA +wrapper for each zigzag half, materializes the full latent KV cache, and routes +the two query halves through the active CP strategy. This mirrors the CP-v2 +FlashAttention path supplied by PR #31619 while preserving FlashInfer's paged +cache format. + ### Positions Position IDs remain CP-sharded after the model-entry split. KDA's forward path @@ -113,8 +136,10 @@ Focused CPU unit tests will use a recording strategy to verify: - Model entry to first KDA gathers `hidden_states` and accepts `residual=None`. - KDA-to-MLA splits both `hidden_states` and `residual`. -- MLA-to-KDA gathers both tensors. -- KDA-to-KDA, MLA-to-MLA, decode, and inactive CP-v2 paths are no-ops. +- MLA attention-to-MLP gathers both tensors. +- KDA-to-KDA, decode, and inactive CP-v2 paths are no-ops. +- KDA recurrent parameters are loaded with global TP rank under CP-v2. +- FlashInfer MLA dispatch materializes latent KV and uses zigzag attention. - Communicator integration uses the model's configured KDA/MLA layer sequence. The existing zigzag strategy tests cover permutation and ragged-batch ordering; diff --git a/python/sglang/srt/layers/cp/kimi_linear.py b/python/sglang/srt/layers/cp/kimi_linear.py index c74f31de6f0d..051da946749c 100644 --- a/python/sglang/srt/layers/cp/kimi_linear.py +++ b/python/sglang/srt/layers/cp/kimi_linear.py @@ -34,9 +34,14 @@ def __init__( *, is_kda_layer: bool, previous_is_kda_layer: Optional[bool], + is_last_layer: bool = False, ) -> None: - self._gather_before_attn = is_kda_layer and (previous_is_kda_layer is not True) - self._shard_before_attn = not is_kda_layer and previous_is_kda_layer is True + self._is_last_layer = is_last_layer + is_first_layer = previous_is_kda_layer is None + # CP-v2 enters the model sharded. Every MLP exits with a full TP batch. + self._gather_before_attn = is_kda_layer and is_first_layer + self._shard_before_attn = not is_kda_layer and not is_first_layer + self._gather_before_mlp = not is_kda_layer def prepare_attn( self, @@ -63,3 +68,41 @@ def prepare_attn( if residual is not None: residual = strategy.shard_hidden_states(residual, forward_batch) return hidden_states, residual + + def prepare_mlp( + self, + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + forward_batch: ForwardBatch, + stream: Optional[Any] = None, + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Gather MLA outputs so normalization and MLP run with global TP.""" + if not self._gather_before_mlp or not is_cp_v2_active(forward_batch): + return hidden_states, residual + + strategy = get_cp_strategy() + assert strategy is not None + hidden_states = strategy.gather_hidden_states( + hidden_states, forward_batch, stream + ) + if residual is not None: + residual = strategy.gather_hidden_states(residual, forward_batch, stream) + return hidden_states, residual + + def postprocess_layer( + self, + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + forward_batch: ForwardBatch, + stream: Optional[Any] = None, + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Shard the full TP output for the model-boundary CP gather.""" + if not self._is_last_layer or not is_cp_v2_active(forward_batch): + return hidden_states, residual + + strategy = get_cp_strategy() + assert strategy is not None + hidden_states = strategy.shard_hidden_states(hidden_states, forward_batch) + if residual is not None: + residual = strategy.shard_hidden_states(residual, forward_batch) + return hidden_states, residual diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index b3ea15f7aee4..4e571e3cd2ce 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -43,7 +43,6 @@ from sglang.srt.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name, - sharded_weight_loader, ) from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA as KimiMLAAttention from sglang.srt.models.llama import LlamaMLP as KimiMLP @@ -53,6 +52,23 @@ from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs +def _global_tp_sharded_weight_loader(shard_axis: int): + """Shard KDA state parameters over the global TP group. + + The generic loader follows attention TP, which is size one under CP-v2. + KDA remains tensor parallel over all TP ranks, so its state parameters must + use the global TP rank as well. + """ + + def loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None: + shard_size = param.data.shape[shard_axis] + start_idx = get_parallel().tp_rank * shard_size + loaded_weight = loaded_weight.narrow(shard_axis, start_idx, shard_size) + default_weight_loader(param, loaded_weight) + + return loader + + class KimiMoE(nn.Module): def __init__( self, @@ -281,7 +297,9 @@ def __init__( torch.empty(divide(projection_size, self.tp_size), dtype=torch.float32) ) - set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)}) + set_weight_attrs( + self.dt_bias, {"weight_loader": _global_tp_sharded_weight_loader(0)} + ) self.qkv_conv1d = MergedColumnParallelLinear( input_size=self.conv_size, @@ -299,7 +317,9 @@ def __init__( self.A_log = nn.Parameter( torch.empty(1, 1, self.local_num_heads, 1, dtype=torch.float32) ) - set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(2)}) + set_weight_attrs( + self.A_log, {"weight_loader": _global_tp_sharded_weight_loader(2)} + ) self.o_norm = FusedRMSNormGated( self.head_dim, eps=rms_norm_eps, activation="sigmoid" @@ -429,6 +449,7 @@ def __init__( previous_is_kda_layer=( config.is_kda_layer(layer_idx - 1) if layer_idx > 0 else None ), + is_last_layer=layer_idx == config.num_hidden_layers - 1, ) if is_kda_layer: @@ -512,9 +533,20 @@ def forward( ) # Fully Connected + hidden_states, residual = self.cp_communicator.prepare_mlp( + hidden_states, + residual, + forward_batch, + torch.cuda.current_stream(), + ) hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) hidden_states = self.mlp(hidden_states) - return hidden_states, residual + return self.cp_communicator.postprocess_layer( + hidden_states, + residual, + forward_batch, + torch.cuda.current_stream(), + ) class KimiLinearModel(nn.Module): diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index c8c19d1ec504..3041b932ea01 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -17,6 +17,7 @@ KimiDeltaAttention, KimiLinearForCausalLM, KimiLinearModel, + _global_tp_sharded_weight_loader, ) from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci @@ -53,6 +54,65 @@ def cp_all_gather_into_tensor_async(self, output, input_tensor, stream): class TestKimiLinearCPV2LayerCommunicator(CustomTestCase): + def test_mla_gathers_hidden_states_and_residual_before_mlp(self): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=False, + previous_is_kda_layer=True, + ) + hidden_states = torch.arange(4).view(2, 2) + residual = hidden_states + 100 + + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, output_residual = communicator.prepare_mlp( + hidden_states, + residual, + SimpleNamespace(), + ) + + self.assertEqual(strategy.gather_calls, 2) + torch.testing.assert_close(output, hidden_states + 10) + torch.testing.assert_close(output_residual, residual + 10) + + def test_last_layer_shards_full_mlp_output_for_model_boundary_gather(self): + strategy = _RecordingStrategy() + communicator = KimiLinearCPV2LayerCommunicator( + is_kda_layer=False, + previous_is_kda_layer=True, + is_last_layer=True, + ) + hidden_states = torch.arange(4).view(2, 2) + residual = hidden_states + 100 + + with ( + patch( + "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", + return_value=True, + ), + patch( + "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", + return_value=strategy, + ), + ): + output, output_residual = communicator.postprocess_layer( + hidden_states, + residual, + SimpleNamespace(), + ) + + self.assertEqual(strategy.shard_calls, 2) + torch.testing.assert_close(output, hidden_states[::2]) + torch.testing.assert_close(output_residual, residual[::2]) + def test_first_kda_layer_gathers_model_entry_shard(self): strategy = _RecordingStrategy() communicator = KimiLinearCPV2LayerCommunicator( @@ -110,7 +170,7 @@ def test_kda_to_mla_shards_hidden_states_and_residual(self): torch.testing.assert_close(output, hidden_states[::2]) torch.testing.assert_close(output_residual, residual[::2]) - def test_mla_to_kda_gathers_hidden_states_and_residual(self): + def test_kda_after_mla_receives_full_mlp_output(self): strategy = _RecordingStrategy() communicator = KimiLinearCPV2LayerCommunicator( is_kda_layer=True, @@ -135,9 +195,10 @@ def test_mla_to_kda_gathers_hidden_states_and_residual(self): SimpleNamespace(), ) - self.assertEqual(strategy.gather_calls, 2) - torch.testing.assert_close(output, hidden_states + 10) - torch.testing.assert_close(output_residual, residual + 10) + self.assertEqual(strategy.gather_calls, 0) + self.assertEqual(strategy.shard_calls, 0) + self.assertIs(output, hidden_states) + self.assertIs(output_residual, residual) def test_same_layout_and_inactive_cp_v2_are_noops(self): hidden_states = torch.arange(8).view(4, 2) @@ -145,7 +206,6 @@ def test_same_layout_and_inactive_cp_v2_are_noops(self): for is_kda_layer, previous_is_kda_layer, cp_v2_active in ( (True, True, True), - (False, False, True), (False, None, True), (True, False, False), (False, True, False), @@ -193,10 +253,6 @@ def test_zigzag_kda_mla_kda_round_trip_restores_token_order(self): is_kda_layer=False, previous_is_kda_layer=True, ) - gather_communicator = KimiLinearCPV2LayerCommunicator( - is_kda_layer=True, - previous_is_kda_layer=False, - ) forward_batches = [] local_hidden_states = [] @@ -258,7 +314,7 @@ def _pad_rank_tensors(rank_tensors): return_value=strategy, ), ): - gathered_hidden, gathered_residual = gather_communicator.prepare_attn( + gathered_hidden, gathered_residual = shard_communicator.prepare_mlp( local_hidden_states[rank], local_residuals[rank], forward_batches[rank], @@ -282,6 +338,7 @@ def test_layer_prepares_cp_layout_before_input_norm(self): intermediate_size=8, hidden_act="silu", rms_norm_eps=1e-5, + num_hidden_layers=5, is_kda_layer=lambda layer_idx: layer_idx != 3, ) communicator = MagicMock() @@ -289,16 +346,26 @@ def test_layer_prepares_cp_layout_before_input_norm(self): prepared_hidden_states = hidden_states + 10 normalized_hidden_states = hidden_states + 20 attention_output = hidden_states + 30 + gathered_attention_output = hidden_states + 35 + gathered_residual = hidden_states + 15 post_norm_output = hidden_states + 40 post_norm_residual = hidden_states + 50 mlp_output = hidden_states + 60 communicator.prepare_attn.return_value = (prepared_hidden_states, None) + communicator.prepare_mlp.return_value = ( + gathered_attention_output, + gathered_residual, + ) input_layernorm = MagicMock(return_value=normalized_hidden_states) self_attn = MagicMock(return_value=attention_output) post_attention_layernorm = MagicMock( return_value=(post_norm_output, post_norm_residual) ) mlp = MagicMock(return_value=mlp_output) + communicator.postprocess_layer.return_value = ( + mlp_output, + post_norm_residual, + ) stream = MagicMock() forward_batch = SimpleNamespace() @@ -334,6 +401,7 @@ def test_layer_prepares_cp_layout_before_input_norm(self): communicator_cls.assert_called_once_with( is_kda_layer=False, previous_is_kda_layer=True, + is_last_layer=False, ) communicator.prepare_attn.assert_called_once_with( hidden_states, @@ -342,6 +410,23 @@ def test_layer_prepares_cp_layout_before_input_norm(self): stream, ) input_layernorm.assert_called_once_with(prepared_hidden_states) + communicator.prepare_mlp.assert_called_once_with( + attention_output, + prepared_hidden_states, + forward_batch, + stream, + ) + post_attention_layernorm.assert_called_once_with( + gathered_attention_output, + gathered_residual, + ) + mlp.assert_called_once_with(post_norm_output) + communicator.postprocess_layer.assert_called_once_with( + mlp_output, + post_norm_residual, + forward_batch, + stream, + ) torch.testing.assert_close(output, mlp_output) torch.testing.assert_close(output_residual, post_norm_residual) @@ -480,6 +565,21 @@ def test_kda_cache_shape_uses_global_tp_size(self): self.assertEqual(shape.conv, [(3, 3072)]) self.assertEqual(shape.num_k_heads_per_tp, 8) + def test_kda_state_weight_loader_uses_global_tp_rank(self): + param = torch.nn.Parameter(torch.zeros(2)) + loaded_weight = torch.arange(8, dtype=param.dtype) + loader = _global_tp_sharded_weight_loader(0) + + with get_parallel().override( + tp_size=4, + tp_rank=2, + attn_tp_size=1, + attn_tp_rank=0, + ): + loader(param, loaded_weight) + + torch.testing.assert_close(param, torch.tensor([4.0, 5.0])) + class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): From 697fff4ded6740d07939cfa67d2ff07c942ad649 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 18 Jul 2026 00:11:57 -0700 Subject: [PATCH 30/35] docs: target merged CP-v2 base --- docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md | 10 +++++----- .../specs/2026-07-17-kimi-linear-cp-v2-design.md | 6 +++--- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md index 12f17ae68d66..ce38d141951e 100644 --- a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md +++ b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md @@ -238,7 +238,7 @@ focused regression test where feasible, and repeat local plus GB300 validation. **Files:** -- Review: all files changed since `origin/cp-v2-mla-prefill`. +- Review: all files changed from the merged PR #31619 implementation on `main`. **Step 1: Re-run verification before claiming completion** @@ -254,8 +254,8 @@ state. Push `codex/kimi-linear-cp-v2` to `Fridge003/sglang`. -**Step 4: Open a stacked PR** +**Step 4: Open the PR** -Open the PR with base `sgl-project:cp-v2-mla-prefill`. Include the layout -transition table, dependency on #31619, unit-test commands, exact GB300 launch -command, and GSM8K result. +PR #31619 merged before publication, so rebase onto current +`sgl-project/sglang:main` and target `main`. Include the layout transition +table, unit-test commands, exact GB300 launch command, and GSM8K result. diff --git a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md index 6a2150ec2976..44bc8f6a1f00 100644 --- a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md +++ b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md @@ -154,6 +154,6 @@ KV-cache errors. ## Delivery -The implementation will be a draft stacked PR targeting -`sgl-project/sglang:cp-v2-mla-prefill`. After PR #31619 merges, the PR can be -retargeted or rebased onto `main`. +PR #31619 merged while this implementation was being verified. The feature +branch is therefore rebased onto its merged CP-v2 implementation and the PR +targets `sgl-project/sglang:main` directly. From c422342e63c2d59163c83c2880e6ca3060a3ea4c Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 18 Jul 2026 00:19:17 -0700 Subject: [PATCH 31/35] fix: address Kimi CP-v2 review feedback --- .../attention/flashinfer_mla_backend.py | 7 +- python/sglang/srt/layers/cp/kimi_linear.py | 8 +- python/sglang/srt/models/kimi_linear.py | 3 - test/registered/cp/test_kimi_linear_cp_v2.py | 106 ++++++++++++++++-- 4 files changed, 109 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 853151ef8547..235835e2fc24 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -388,6 +388,11 @@ def init_forward_metadata_out_graph( def init_forward_metadata(self, forward_batch: ForwardBatch): self.cp_prefill_metadata = None + if is_cp_v2_active(forward_batch): + # CP wrappers are planned lazily after the eager runner builds the + # per-rank zigzag metadata. The ordinary full-batch plan is unused. + self.forward_metadata = None + return if forward_batch.forward_mode.is_decode_or_idle(): self.indices_updater_decode.update( forward_batch.req_pool_indices, @@ -565,7 +570,7 @@ def _plan_cp_prefill_wrapper( updater.num_local_heads, updater.kv_lora_rank, updater.qk_rope_head_dim, - self.page_size, + 1, True, updater.scaling, updater.q_data_type, diff --git a/python/sglang/srt/layers/cp/kimi_linear.py b/python/sglang/srt/layers/cp/kimi_linear.py index 051da946749c..9cbc602748ad 100644 --- a/python/sglang/srt/layers/cp/kimi_linear.py +++ b/python/sglang/srt/layers/cp/kimi_linear.py @@ -18,11 +18,11 @@ from typing import TYPE_CHECKING, Any, Optional, Tuple +import torch + from sglang.srt.layers.cp.utils import get_cp_strategy, is_cp_v2_active if TYPE_CHECKING: - import torch - from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -56,6 +56,8 @@ def prepare_attn( strategy = get_cp_strategy() assert strategy is not None if self._gather_before_attn: + if stream is None: + stream = torch.cuda.current_stream() hidden_states = strategy.gather_hidden_states( hidden_states, forward_batch, stream ) @@ -82,6 +84,8 @@ def prepare_mlp( strategy = get_cp_strategy() assert strategy is not None + if stream is None: + stream = torch.cuda.current_stream() hidden_states = strategy.gather_hidden_states( hidden_states, forward_batch, stream ) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 4e571e3cd2ce..660aae1e6266 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -515,7 +515,6 @@ def forward( hidden_states, residual, forward_batch, - torch.cuda.current_stream(), ) # Self Attention @@ -537,7 +536,6 @@ def forward( hidden_states, residual, forward_batch, - torch.cuda.current_stream(), ) hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) hidden_states = self.mlp(hidden_states) @@ -545,7 +543,6 @@ def forward( hidden_states, residual, forward_batch, - torch.cuda.current_stream(), ) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index 3041b932ea01..f193ab057410 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -29,11 +29,13 @@ class _RecordingStrategy: def __init__(self): self.gather_calls = 0 + self.gather_streams = [] self.shard_calls = 0 def gather_hidden_states(self, hidden_states, forward_batch, stream=None): - del forward_batch, stream + del forward_batch self.gather_calls += 1 + self.gather_streams.append(stream) return hidden_states + 10 def shard_hidden_states(self, hidden_states, forward_batch): @@ -63,6 +65,7 @@ def test_mla_gathers_hidden_states_and_residual_before_mlp(self): hidden_states = torch.arange(4).view(2, 2) residual = hidden_states + 100 + stream = MagicMock() with ( patch( "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", @@ -72,6 +75,7 @@ def test_mla_gathers_hidden_states_and_residual_before_mlp(self): "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", return_value=strategy, ), + patch("torch.cuda.current_stream", return_value=stream), ): output, output_residual = communicator.prepare_mlp( hidden_states, @@ -80,6 +84,7 @@ def test_mla_gathers_hidden_states_and_residual_before_mlp(self): ) self.assertEqual(strategy.gather_calls, 2) + self.assertEqual(strategy.gather_streams, [stream, stream]) torch.testing.assert_close(output, hidden_states + 10) torch.testing.assert_close(output_residual, residual + 10) @@ -121,6 +126,7 @@ def test_first_kda_layer_gathers_model_entry_shard(self): ) hidden_states = torch.arange(4).view(2, 2) + stream = MagicMock() with ( patch( "sglang.srt.layers.cp.kimi_linear.is_cp_v2_active", @@ -130,6 +136,7 @@ def test_first_kda_layer_gathers_model_entry_shard(self): "sglang.srt.layers.cp.kimi_linear.get_cp_strategy", return_value=strategy, ), + patch("torch.cuda.current_stream", return_value=stream), ): output, residual = communicator.prepare_attn( hidden_states, @@ -138,6 +145,7 @@ def test_first_kda_layer_gathers_model_entry_shard(self): ) self.assertEqual(strategy.gather_calls, 1) + self.assertEqual(strategy.gather_streams, [stream]) torch.testing.assert_close(output, hidden_states + 10) self.assertIsNone(residual) @@ -366,7 +374,6 @@ def test_layer_prepares_cp_layout_before_input_norm(self): mlp_output, post_norm_residual, ) - stream = MagicMock() forward_batch = SimpleNamespace() with ( @@ -387,7 +394,6 @@ def test_layer_prepares_cp_layout_before_input_norm(self): "sglang.srt.models.kimi_linear.RMSNorm", side_effect=[input_layernorm, post_attention_layernorm], ), - patch("torch.cuda.current_stream", return_value=stream), ): layer = KimiDecoderLayer(config=config, layer_idx=3) output, output_residual = layer( @@ -407,14 +413,12 @@ def test_layer_prepares_cp_layout_before_input_norm(self): hidden_states, None, forward_batch, - stream, ) input_layernorm.assert_called_once_with(prepared_hidden_states) communicator.prepare_mlp.assert_called_once_with( attention_output, prepared_hidden_states, forward_batch, - stream, ) post_attention_layernorm.assert_called_once_with( gathered_attention_output, @@ -425,7 +429,6 @@ def test_layer_prepares_cp_layout_before_input_norm(self): mlp_output, post_norm_residual, forward_batch, - stream, ) torch.testing.assert_close(output, mlp_output) torch.testing.assert_close(output_residual, post_norm_residual) @@ -603,6 +606,70 @@ def test_inner_model_accepts_cp_v2_input_embeds_keyword(self): class TestKimiLinearFlashInferMLACP(CustomTestCase): + def test_cp_v2_skips_unused_full_batch_prefill_plan(self): + backend = object.__new__(FlashInferMLAAttnBackend) + backend.cp_prefill_metadata = MagicMock() + backend.forward_metadata = MagicMock() + backend.indices_updater_decode = MagicMock() + backend.indices_updater_prefill = MagicMock() + forward_batch = SimpleNamespace( + forward_mode=SimpleNamespace( + is_decode_or_idle=lambda: False, + is_target_verify=lambda: False, + ) + ) + + with patch( + "sglang.srt.layers.attention.flashinfer_mla_backend.is_cp_v2_active", + return_value=True, + ): + backend.init_forward_metadata(forward_batch) + + self.assertIsNone(backend.cp_prefill_metadata) + self.assertIsNone(backend.forward_metadata) + backend.indices_updater_decode.update.assert_not_called() + backend.indices_updater_prefill.update.assert_not_called() + + def test_cp_wrapper_plan_uses_physical_token_page_size(self): + backend = object.__new__(FlashInferMLAAttnBackend) + backend.workspace_buffer = MagicMock() + backend.page_size = 16 + backend.req_to_token_pool = SimpleNamespace( + req_to_token=torch.zeros((2, 32), dtype=torch.int32) + ) + backend.indices_updater_prefill = SimpleNamespace( + num_local_heads=4, + kv_lora_rank=8, + qk_rope_head_dim=4, + scaling=0.5, + q_data_type=torch.bfloat16, + data_type=torch.bfloat16, + ) + forward_batch = SimpleNamespace(req_pool_indices=torch.tensor([0, 1])) + qo_indptr = torch.tensor([0, 2, 5], dtype=torch.int32) + kv_lens = torch.tensor([3, 4], dtype=torch.int32) + wrapper = MagicMock() + indices_kernel = MagicMock() + + with ( + patch( + "sglang.srt.layers.attention.flashinfer_mla_backend.BatchMLAPagedAttentionWrapper", + return_value=wrapper, + ), + patch( + "sglang.srt.layers.attention.flashinfer_mla_backend.create_flashinfer_kv_indices_triton", + indices_kernel, + ), + ): + backend._plan_cp_prefill_wrapper( + forward_batch, + qo_indptr, + kv_lens, + kv_lens_sum=7, + ) + + self.assertEqual(wrapper.plan.call_args.args[7], 1) + def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): backend = object.__new__(FlashInferMLAAttnBackend) backend.device = torch.device("cpu") @@ -611,8 +678,17 @@ def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): ) backend._run_cp_paged_attention = MagicMock() strategy = MagicMock() - expected = torch.randn(5, 2, 4) - strategy.run_attention.return_value = expected + + def run_attention(q_fused, forward_batch, device, attn_fn, **kwargs): + del forward_batch, device, kwargs + return torch.cat( + [ + attn_fn(q_fused[:2], None, None, None), + attn_fn(q_fused[2:], None, None, None), + ] + ) + + strategy.run_attention.side_effect = run_attention q = torch.randn(5, 8) q_rope = torch.randn(5, 4) k = torch.randn(5, 8) @@ -624,6 +700,9 @@ def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): head_dim=6, ) forward_batch = SimpleNamespace() + backend._run_cp_paged_attention.side_effect = ( + lambda wrapper, q_chunk, layer: q_chunk[..., : layer.v_head_dim] + ) with ( patch( @@ -652,7 +731,16 @@ def test_cp_v2_materializes_full_latent_and_dispatches_zigzag_attention(self): k_rope, ) strategy.run_attention.assert_called_once() - torch.testing.assert_close(output, expected.view(5, 8)) + self.assertEqual(backend._run_cp_paged_attention.call_count, 2) + self.assertIs( + backend._run_cp_paged_attention.call_args_list[0].args[0], + backend._get_cp_prefill_metadata.return_value.wrappers[0], + ) + self.assertIs( + backend._run_cp_paged_attention.call_args_list[1].args[0], + backend._get_cp_prefill_metadata.return_value.wrappers[1], + ) + torch.testing.assert_close(output, q) if __name__ == "__main__": From 8e9d30629bba2ebf2c3c3173867939b8e9790eae Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sat, 18 Jul 2026 00:31:29 -0700 Subject: [PATCH 32/35] docs: remove implementation planning files --- .../plans/2026-07-17-kimi-linear-cp-v2.md | 261 ------------------ .../2026-07-17-kimi-linear-cp-v2-design.md | 159 ----------- 2 files changed, 420 deletions(-) delete mode 100644 docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md delete mode 100644 docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md diff --git a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md b/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md deleted file mode 100644 index ce38d141951e..000000000000 --- a/docs/superpowers/plans/2026-07-17-kimi-linear-cp-v2.md +++ /dev/null @@ -1,261 +0,0 @@ -# Kimi-Linear CP-v2 Implementation Plan - -> **For Codex:** Execute this plan test-first in the current branch. Keep the -> production change limited to the Kimi-Linear layer transitions and CP-v2 -> activation hooks. - -**Goal:** Run Kimi-Linear MLA layers on zigzag CP shards while running KDA layers -on complete token batches replicated across the four TP ranks. - -**Architecture:** A Kimi-specific CP-v2 layer communicator converts -`hidden_states` and `residual` at KDA/MLA boundaries using the active -`ContextParallelStrategy`. `KimiDecoderLayer` invokes it before input RMSNorm -and after MLA attention, before the TP MLP. The existing CP-v2 eager runner -continues to own model-entry splitting and model-exit gathering. - -**Tech Stack:** Python, PyTorch distributed collectives, SGLang CP-v2 strategy -API, `unittest`, four NVIDIA GB300 GPUs, GSM8K evaluation. - ---- - -### Task 1: Specify layer-transition behavior with a failing unit test - -**Files:** - -- Create: `test/registered/cp/test_kimi_linear_cp_v2.py` -- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` - -**Step 1: Write the first failing test** - -Create a CPU-registered `CustomTestCase` with a recording CP strategy. The first -test constructs a communicator for the first KDA layer, patches CP-v2 active, -and asserts that its rank-local `hidden_states` are passed through -`gather_hidden_states` while `residual=None` is preserved. - -```python -communicator = KimiLinearCPV2LayerCommunicator( - is_kda_layer=True, - previous_is_kda_layer=None, -) -hidden_states, residual = communicator.prepare_attn( - hidden_states, None, forward_batch -) -self.assertEqual(strategy.gather_calls, 1) -self.assertIsNone(residual) -``` - -**Step 2: Run the focused test and confirm RED** - -Run: - -```bash -python -m unittest test.registered.cp.test_kimi_linear_cp_v2 -v -``` - -Expected: failure because the communicator module/class does not exist. - -### Task 2: Implement the minimal CP-v2 communicator - -**Files:** - -- Create: `python/sglang/srt/layers/cp/kimi_linear.py` -- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` -- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` - -**Step 1: Implement first-KDA gather** - -Define `KimiLinearCPV2LayerCommunicator` with static layer-type inputs and a -`prepare_attn` method. Gate it with `is_cp_v2_active`; obtain the strategy with -`get_cp_strategy`; gather the model-entry shard for a first KDA layer. - -**Step 2: Run the focused test and confirm GREEN** - -Run the same unit-test command and require it to pass. - -**Step 3: Add one failing transition test at a time** - -Add and run tests for: - -- KDA to MLA: shard `hidden_states` and `residual`. -- MLA attention to MLP: gather `hidden_states` and `residual`. -- KDA to KDA: identity/no strategy calls. -- CP-v2 inactive: identity/no strategy calls. - -For each case, first observe failure, then implement the smallest transition -logic needed to pass it. - -**Step 4: Add a real zigzag round-trip test** - -Use `ZigzagCPStrategy` metadata and the existing fake CP group pattern to prove -that a full tensor split for MLA and gathered before its MLP returns to original -token order, including its residual tensor. - -**Step 5: Run communicator and existing strategy tests** - -```bash -python -m unittest \ - test.registered.cp.test_kimi_linear_cp_v2 \ - test.registered.cp.test_cp_strategy_unit -v -``` - -Expected: all pass. - -### Task 3: Wire the communicator into Kimi-Linear - -**Files:** - -- Modify: `python/sglang/srt/models/kimi_linear.py` -- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` -- Test: `test/registered/cp/test_kimi_linear_cp_v2.py` - -**Step 1: Add a failing wiring test** - -Verify the model exposes a communicator configured from each layer's current -and previous KDA/MLA type, including `previous_is_kda_layer=None` for layer zero. -Use a minimal Kimi config or constructor patching so the test remains CPU-only. - -**Step 2: Construct and invoke the communicator** - -In `KimiDecoderLayer.__init__`, compute current and previous layer types from -`KimiLinearConfig.is_kda_layer`. Construct the communicator. At the very start -of `forward`, call `prepare_attn` before input RMSNorm. - -After MLA attention, call `prepare_mlp` before post-attention RMSNorm and the -TP MLP. Shard the final full layer output so the generic model-exit gather keeps -its CP-v2 contract. - -**Step 3: Run the focused test and confirm GREEN** - -```bash -python -m unittest test.registered.cp.test_kimi_linear_cp_v2 -v -``` - -### Task 4: Enable Kimi-Linear in the CP-v2 eager path - -**Files:** - -- Modify: `python/sglang/srt/layers/cp/utils.py` -- Modify: `python/sglang/srt/models/kimi_linear.py` -- Modify: `test/registered/cp/test_kimi_linear_cp_v2.py` - -**Step 1: Add failing activation/accessor assertions** - -Assert that `KimiLinearForCausalLM` is in `CP_V2_DEFAULT_MODEL_CLASSES` and -that its `get_input_embeddings()` accessor returns `model.embed_tokens`. - -**Step 2: Add the activation and embedding hooks** - -Add the architecture string to the default class set and the standard accessor -to `KimiLinearForCausalLM`. - -**Step 3: Run the focused CP tests** - -```bash -python -m unittest \ - test.registered.cp.test_kimi_linear_cp_v2 \ - test.registered.cp.test_cp_strategy_unit -v -``` - -Expected: all pass. - -### Task 4b: Keep KDA on global tensor parallelism - -Use global TP rank/size for KDA projections, recurrent cache shape, head -partitioning, and `A_log`/`dt_bias` weight loading. Add a regression test where -`tp_rank` differs from `attn_tp_rank`. - -Add the FlashInfer MLA CP-v2 path because it is the GB300 default backend, and -test its latent-KV materialization plus zigzag dispatch. - -### Task 5: Local static and regression verification - -**Files:** - -- Verify all modified Python and test files. - -**Step 1: Run formatting and lint checks** - -```bash -pre-commit run --files \ - python/sglang/srt/layers/cp/kimi_linear.py \ - python/sglang/srt/layers/cp/utils.py \ - python/sglang/srt/models/kimi_linear.py \ - test/registered/cp/test_kimi_linear_cp_v2.py -``` - -**Step 2: Run CP-focused test suites** - -Run the new unit test, existing CP strategy unit test, and any applicable -server-argument tests selected by the diff. - -**Step 3: Inspect the final local diff** - -Require `git diff --check`, review all changes against the design, and confirm -no unrelated user changes are present. - -### Task 6: GB300 end-to-end verification - -**Files:** - -- Remote checkout and logs on `baizhou-dev-2`. -- No generated benchmark artifacts committed to the repository. - -**Step 1: Prepare the devbox** - -Use `rx devbox run baizhou-dev-2`. Clone or update SGLang, fetch the implementation -branch, pull the latest stacked base, and install the current editable Python and -kernel dependencies before launching a job. - -**Step 2: Locate or download the model** - -Use `moonshotai/Kimi-Linear-48B-A3B-Instruct` from a shared cache if present; -otherwise download it to devbox-attached persistent storage. - -**Step 3: Launch the requested configuration** - -Start SGLang with: - -```bash ---tp 4 --attn-cp-size 4 --enable-prefill-cp --cp-strategy zigzag -``` - -Capture the exact command, commit, model path, backend selection, and complete -server log. - -**Step 4: Run GSM8K** - -Use the repository evaluation command against the live endpoint. Record sample -count, accuracy, and any mismatch/error output. If no established Kimi-Linear -threshold exists, compare against a TP4 non-CP run using the same model, -tokenizer, prompts, and decoding settings. - -**Step 5: Diagnose until verified** - -If the server or evaluation fails, preserve the first failure signature, add a -focused regression test where feasible, and repeat local plus GB300 validation. - -### Task 7: Review, commit, push, and open the stacked PR - -**Files:** - -- Review: all files changed from the merged PR #31619 implementation on `main`. - -**Step 1: Re-run verification before claiming completion** - -Capture fresh output for the focused tests, pre-commit checks, and GSM8K result. - -**Step 2: Commit intentionally** - -Keep the design document commit and create logically scoped implementation/test -commits. Do not fold in changes from the stacked base or unrelated worktree -state. - -**Step 3: Push to the authenticated fork** - -Push `codex/kimi-linear-cp-v2` to `Fridge003/sglang`. - -**Step 4: Open the PR** - -PR #31619 merged before publication, so rebase onto current -`sgl-project/sglang:main` and target `main`. Include the layout transition -table, unit-test commands, exact GB300 launch command, and GSM8K result. diff --git a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md b/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md deleted file mode 100644 index 44bc8f6a1f00..000000000000 --- a/docs/superpowers/specs/2026-07-17-kimi-linear-cp-v2-design.md +++ /dev/null @@ -1,159 +0,0 @@ -# Kimi-Linear CP-v2 Layer Transitions - -## Context - -PR #31619 moves MLA prefill context parallelism to the CP-v2 strategy API. At -the model boundary, CP-v2 splits embeddings and positions into the selected CP -layout, and after the model body it gathers hidden states before logits. - -Kimi-Linear alternates two attention implementations with different token -layout requirements: - -- Kimi Delta Attention (KDA) is tensor parallel and must receive the complete - token batch on every TP rank. -- Multi-head Latent Attention (MLA) participates in prefill context parallelism - and must receive the rank-local zigzag token shard. - -The model carries both `hidden_states` and a deferred `residual` between decoder -layers. They must always use the same token layout. - -## Goals - -- Support Kimi-Linear prefill with `--tp 4 --attn-cp-size 4 - --enable-prefill-cp --cp-strategy zigzag`. -- Gather CP-sharded MLA attention output before its post-attention norm and MLP. -- Split replicated layer state before each MLA region. -- Reuse the active CP strategy for ordering, padding, and collectives. -- Preserve the non-CP and decode paths. -- Verify correctness with focused unit tests and GSM8K on four GB300 GPUs. - -## Non-goals - -- Adding a new CP strategy. -- Enabling the unfinished interleave strategy. -- Changing KDA kernels, MLA kernels, or KV-cache semantics. -- Supporting Kimi K2.5's multimodal wrapper in this change. -- Optimizing transition collectives beyond the minimum correct implementation. - -## Design - -### Communicator - -Add `KimiLinearCPV2LayerCommunicator` under -`python/sglang/srt/layers/cp/kimi_linear.py`. The class owns only the transition -between the two token layouts; it does not replace the general TP/DP/MoE -`LayerCommunicator` in `layers/communicator.py`. - -Each `KimiDecoderLayer` constructs the communicator with: - -- Whether the current layer is KDA. -- Whether the preceding global layer is KDA. Layer zero has no preceding layer. -- Whether it is the final decoder layer. - -At the beginning of `KimiDecoderLayer.forward`, the communicator receives -`hidden_states`, `residual`, and `forward_batch`, and returns state in the layout -required by the current attention layer. After attention, MLA output and its -residual are gathered before the post-attention norm and MLP. Consequently, -every MLP runs on the complete token batch and every layer exits its MLP in the -replicated TP layout. The final layer shards that output once more so the -generic CP-v2 model-exit gather retains its normal contract. - -It is active only when `is_cp_v2_active(forward_batch)` is true. Otherwise it -returns its inputs without communication. - -### Transition table - -| Boundary | Incoming state | Operation | -| --- | --- | --- | -| Model entry → first KDA | CP shard | Gather to complete token order | -| MLP → KDA | Replicated | No-op | -| MLP → MLA | Replicated | Split with the active CP strategy | -| MLA attention → MLP | CP shard | Gather to complete token order | -| Final MLP → model exit | Replicated | Split for the existing model-exit gather | - -The first Kimi-Linear layer is KDA, so model-entry embeddings are gathered -before layer zero. Gathering immediately after each MLA attention operation is -required because the following MoE/MLP contains TP all-reduces whose token -dimensions must be identical on all ranks. - -### Hidden state and residual - -When `residual` is present, both tensors undergo the same transition. A gather -uses `ContextParallelStrategy.gather_hidden_states` so zigzag rank ordering, -ragged batches, and padding are restored exactly as at the model boundary. A -split uses `ContextParallelStrategy.shard_hidden_states`. - -The initial layer has `residual=None`; only `hidden_states` is gathered there. -No residual addition is moved across the transition. - -### KDA tensor parallel state - -KDA remains tensor parallel over the global TP group while MLA attention uses -the CP-sharded layout. Its projections, recurrent-cache head shape, and -`A_log`/`dt_bias` state-parameter loading therefore use global `tp_size` and -`tp_rank`, not attention TP. Under `--tp 4 --attn-cp-size 4`, attention TP has -size one; using its rank would incorrectly load rank zero's recurrent -parameters on all four KDA ranks. - -### FlashInfer MLA - -GB300 selects FlashInfer as the default MLA backend. CP-v2 plans one paged MLA -wrapper for each zigzag half, materializes the full latent KV cache, and routes -the two query halves through the active CP strategy. This mirrors the CP-v2 -FlashAttention path supplied by PR #31619 while preserving FlashInfer's paged -cache format. - -### Positions - -Position IDs remain CP-sharded after the model-entry split. KDA's forward path -does not consume `positions`, while MLA requires positions aligned with its -rank-local hidden-state shard. Consequently, position IDs do not need gather -and split transitions. - -### CP-v2 activation - -Add `KimiLinearForCausalLM` to `CP_V2_DEFAULT_MODEL_CLASSES`. Also expose the -model's input embedding layer through `get_input_embeddings`, which the CP-v2 -eager runner uses to embed the complete token batch before its initial split. - -The change remains restricted to CP-v2 context-parallel extend batches through -the existing `is_cp_v2_active` gate. Decode and legacy CP-v1 behavior are -unchanged. - -## Error handling and invariants - -- The communicator requires an initialized CP strategy whenever CP-v2 is - active; this follows the existing CP-v2 invariant. -- Both tensors must have matching token dimensions before a joint transition. -- The strategy metadata prepared by the eager runner is the single source of - truth for split and gather ordering. -- A transition is selected from static model layer types, not inferred from - runtime tensor lengths. - -## Testing - -Focused CPU unit tests will use a recording strategy to verify: - -- Model entry to first KDA gathers `hidden_states` and accepts `residual=None`. -- KDA-to-MLA splits both `hidden_states` and `residual`. -- MLA attention-to-MLP gathers both tensors. -- KDA-to-KDA, decode, and inactive CP-v2 paths are no-ops. -- KDA recurrent parameters are loaded with global TP rank under CP-v2. -- FlashInfer MLA dispatch materializes latent KV and uses zigzag attention. -- Communicator integration uses the model's configured KDA/MLA layer sequence. - -The existing zigzag strategy tests cover permutation and ragged-batch ordering; -an additional transition round-trip test will ensure the communicator uses -those strategy operations without changing order. - -End-to-end verification will launch Kimi-Linear on `baizhou-dev-2` with four -GB300 GPUs and the requested flags, then run the repository's GSM8K evaluation. -The result will be compared with the same model's non-CP baseline or its known -expected accuracy, and server logs will be checked for collective, shape, and -KV-cache errors. - -## Delivery - -PR #31619 merged while this implementation was being verified. The feature -branch is therefore rebased onto its merged CP-v2 implementation and the PR -targets `sgl-project/sglang:main` directly. From c0f552608019331ef2241b166fe2c8c01b2ff57d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sun, 19 Jul 2026 03:47:53 -0700 Subject: [PATCH 33/35] feat: shard KDA heads over CP ranks --- python/sglang/srt/configs/kimi_linear.py | 20 +- python/sglang/srt/models/kimi_linear.py | 72 +++-- test/registered/cp/test_kimi_linear_cp_v2.py | 316 +++++++++++++++++-- 3 files changed, 360 insertions(+), 48 deletions(-) diff --git a/python/sglang/srt/configs/kimi_linear.py b/python/sglang/srt/configs/kimi_linear.py index a26c947bd0d7..9521a8944acd 100644 --- a/python/sglang/srt/configs/kimi_linear.py +++ b/python/sglang/srt/configs/kimi_linear.py @@ -7,6 +7,19 @@ from sglang.srt.runtime_context import get_parallel +def _get_kda_head_shard_info() -> tuple[int, int, bool]: + """Return the rank and size that own KDA heads. + + CP ranks own KDA heads when attention CP is configured. Without CP, KDA + keeps its existing global-TP ownership. + """ + + parallel = get_parallel() + if parallel.attn_cp_size > 1: + return parallel.attn_cp_rank, parallel.attn_cp_size, True + return parallel.tp_rank, parallel.tp_size, False + + class KimiLinearConfig(PretrainedConfig): model_type = "kimi_linear" keys_to_ignore_at_inference = ["past_key_values"] @@ -152,9 +165,12 @@ def full_attention_layer_ids(self): @property def mamba2_cache_params(self) -> KimiLinearCacheParams: - + parallel = get_parallel() + kda_head_shard_size = ( + parallel.attn_cp_size if parallel.attn_cp_size > 1 else parallel.tp_size + ) shape = KimiLinearStateShape.create( - tp_world_size=get_parallel().tp_size, + tp_world_size=kda_head_shard_size, num_heads=self.linear_attn_config["num_heads"], head_dim=self.linear_attn_config["head_dim"], conv_kernel_size=self.linear_attn_config["short_conv_kernel_size"], diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 660aae1e6266..b4fb51af1117 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -9,7 +9,10 @@ from torch import nn from sglang.kernels.ops.attention.fla.fused_norm_gate import FusedRMSNormGated -from sglang.srt.configs.kimi_linear import KimiLinearConfig +from sglang.srt.configs.kimi_linear import ( + KimiLinearConfig, + _get_kda_head_shard_info, +) from sglang.srt.distributed import ( divide, get_pp_group, @@ -52,17 +55,13 @@ from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs -def _global_tp_sharded_weight_loader(shard_axis: int): - """Shard KDA state parameters over the global TP group. - - The generic loader follows attention TP, which is size one under CP-v2. - KDA remains tensor parallel over all TP ranks, so its state parameters must - use the global TP rank as well. - """ +def _kda_head_sharded_weight_loader(shard_axis: int): + """Shard KDA state parameters over their head-ownership group.""" def loader(param: torch.Tensor, loaded_weight: torch.Tensor) -> None: shard_size = param.data.shape[shard_axis] - start_idx = get_parallel().tp_rank * shard_size + head_shard_rank, _, _ = _get_kda_head_shard_info() + start_idx = head_shard_rank * shard_size loaded_weight = loaded_weight.narrow(shard_axis, start_idx, shard_size) default_weight_loader(param, loaded_weight) @@ -194,6 +193,11 @@ def __init__( ) -> None: super().__init__() self.tp_size = get_parallel().tp_size + ( + self.head_shard_rank, + self.head_shard_size, + self.kda_heads_use_cp, + ) = _get_kda_head_shard_info() self.hidden_size = hidden_size self.config = config self.head_dim = config.linear_attn_config["head_dim"] @@ -204,14 +208,17 @@ def __init__( self.head_v_dim = config.v_head_dim self.layer_idx = layer_idx self.prefix = prefix - assert self.num_heads % self.tp_size == 0 - self.local_num_heads = divide(self.num_heads, self.tp_size) + assert self.num_heads % self.head_shard_size == 0 + self.local_num_heads = divide(self.num_heads, self.head_shard_size) projection_size = self.head_dim * self.num_heads self.conv_size = config.linear_attn_config["short_conv_kernel_size"] - # TODO: support fusion with quant - self.do_fuse_qkvbfg = quant_config is None + # The fused projections currently derive their shard from global TP. + # Use them only when that is also the KDA head-ownership group. + self.do_fuse_qkvbfg = ( + quant_config is None and self.head_shard_size == self.tp_size + ) if self.do_fuse_qkvbfg: # Fuse: q, k, v, beta (column parallel) + f_a, g_a (replicated) @@ -231,8 +238,8 @@ def __init__( prefix=f"{prefix}.fused_qkvbfg_a_proj", ) self.split_sizes = [ - 3 * projection_size // self.tp_size, # qkv - self.num_heads // self.tp_size, # beta + 3 * projection_size // self.head_shard_size, # qkv + self.num_heads // self.head_shard_size, # beta 2 * self.head_dim, # f_a, g_a ] self.fused_fg_b_proj = ColumnParallelBatchedLinear( @@ -240,7 +247,6 @@ def __init__( ) else: # Unfused path: separate QKVParallelLinear - tp_rank = get_parallel().tp_rank self.qkv_proj = QKVParallelLinear( self.hidden_size, self.head_dim, @@ -248,8 +254,8 @@ def __init__( self.num_k_heads, bias=False, quant_config=quant_config, - tp_rank=tp_rank, - tp_size=self.tp_size, + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, v_head_size=self.head_v_dim, prefix=f"{prefix}.qkv_proj", ) @@ -268,6 +274,8 @@ def __init__( bias=False, quant_config=quant_config, prefix=f"{prefix}.f_b_proj", + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, ) self.b_proj = ColumnParallelLinear( @@ -276,6 +284,8 @@ def __init__( bias=False, quant_config=quant_config, prefix=f"{prefix}.b_proj", + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, ) self.g_a_proj = ReplicatedLinear( @@ -291,14 +301,18 @@ def __init__( bias=False, quant_config=quant_config, prefix=f"{prefix}.g_b_proj", + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, ) self.dt_bias = nn.Parameter( - torch.empty(divide(projection_size, self.tp_size), dtype=torch.float32) + torch.empty( + divide(projection_size, self.head_shard_size), dtype=torch.float32 + ) ) set_weight_attrs( - self.dt_bias, {"weight_loader": _global_tp_sharded_weight_loader(0)} + self.dt_bias, {"weight_loader": _kda_head_sharded_weight_loader(0)} ) self.qkv_conv1d = MergedColumnParallelLinear( @@ -307,6 +321,8 @@ def __init__( bias=False, params_dtype=torch.float32, prefix=f"{prefix}.qkv_conv1d", + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, ) # unsqueeze to fit conv1d weights shape into the linear weights shape. # Can't do this in `weight_loader` since it already exists in @@ -318,7 +334,7 @@ def __init__( torch.empty(1, 1, self.local_num_heads, 1, dtype=torch.float32) ) set_weight_attrs( - self.A_log, {"weight_loader": _global_tp_sharded_weight_loader(2)} + self.A_log, {"weight_loader": _kda_head_sharded_weight_loader(2)} ) self.o_norm = FusedRMSNormGated( @@ -330,6 +346,9 @@ def __init__( bias=False, quant_config=quant_config, prefix=f"{prefix}.o_proj", + tp_rank=self.head_shard_rank, + tp_size=self.head_shard_size, + reduce_results=not self.kda_heads_use_cp, ) conv_weights = self.qkv_conv1d.weight.squeeze(1) @@ -337,9 +356,9 @@ def __init__( self.attn = RadixLinearAttention( layer_id=self.layer_idx, - num_q_heads=self.num_k_heads // self.tp_size, - num_k_heads=self.num_k_heads // self.tp_size, - num_v_heads=self.num_v_heads // self.tp_size, + num_q_heads=self.num_k_heads // self.head_shard_size, + num_k_heads=self.num_k_heads // self.head_shard_size, + num_v_heads=self.num_v_heads // self.head_shard_size, head_q_dim=self.head_k_dim, head_k_dim=self.head_k_dim, head_v_dim=self.head_v_dim, @@ -426,7 +445,10 @@ def forward( core_attn_out = self.o_norm(core_attn_out, norm_gate) core_attn_out = core_attn_out.squeeze(0).flatten(-2) # 1 n h d -> n (h d) - return self.o_proj(core_attn_out)[0] + output = self.o_proj(core_attn_out)[0] + if self.kda_heads_use_cp: + output = get_parallel().attn_cp_group.all_reduce(output) + return output class KimiDecoderLayer(nn.Module): diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index f193ab057410..cf15f16122f6 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -17,7 +17,7 @@ KimiDeltaAttention, KimiLinearForCausalLM, KimiLinearModel, - _global_tp_sharded_weight_loader, + _kda_head_sharded_weight_loader, ) from sglang.srt.runtime_context import get_parallel from sglang.test.ci.ci_register import register_cpu_ci @@ -433,7 +433,7 @@ def test_layer_prepares_cp_layout_before_input_norm(self): torch.testing.assert_close(output, mlp_output) torch.testing.assert_close(output_residual, post_norm_residual) - def test_kda_backend_partitions_heads_over_global_tp(self): + def test_kda_backend_partitions_heads_over_cp(self): config = SimpleNamespace( linear_attn_config={ "head_dim": 128, @@ -447,17 +447,23 @@ def test_kda_backend_partitions_heads_over_global_tp(self): with ( get_parallel().override( - tp_size=4, - tp_rank=2, - attn_tp_size=1, - attn_tp_rank=0, + tp_size=8, + tp_rank=6, + attn_tp_size=2, + attn_tp_rank=1, + attn_cp_size=4, + attn_cp_rank=2, ), patch( - "sglang.srt.models.kimi_linear.MergedColumnParallelRepeatedLinear", + "sglang.srt.models.kimi_linear.QKVParallelLinear", return_value=MagicMock(), ), patch( - "sglang.srt.models.kimi_linear.ColumnParallelBatchedLinear", + "sglang.srt.models.kimi_linear.ReplicatedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.ColumnParallelLinear", return_value=MagicMock(), ), patch( @@ -481,6 +487,7 @@ def test_kda_backend_partitions_heads_over_global_tp(self): layer_idx=0, hidden_size=256, config=config, + quant_config=MagicMock(), ) radix_linear_attention_cls.assert_called_once() @@ -489,7 +496,7 @@ def test_kda_backend_partitions_heads_over_global_tp(self): self.assertEqual(backend_args["num_k_heads"], 8) self.assertEqual(backend_args["num_v_heads"], 8) - def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): + def test_unfused_kda_projections_use_cp_rank_and_size(self): config = SimpleNamespace( linear_attn_config={ "head_dim": 128, @@ -500,6 +507,149 @@ def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): dtype=torch.bfloat16, ) qkv_parallel_linear = MagicMock() + column_parallel_linear = MagicMock() + merged_column_parallel_linear = MagicMock() + row_parallel_linear = MagicMock() + + with ( + get_parallel().override( + tp_size=8, + tp_rank=6, + attn_tp_size=2, + attn_tp_rank=1, + attn_cp_size=4, + attn_cp_rank=2, + ), + patch( + "sglang.srt.models.kimi_linear.QKVParallelLinear", + return_value=qkv_parallel_linear, + ) as qkv_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.ReplicatedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.ColumnParallelLinear", + return_value=column_parallel_linear, + ) as column_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelLinear", + return_value=merged_column_parallel_linear, + ) as merged_column_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.FusedRMSNormGated", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RowParallelLinear", + return_value=row_parallel_linear, + ) as row_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.RadixLinearAttention", + return_value=MagicMock(), + ), + ): + KimiDeltaAttention( + layer_idx=0, + hidden_size=256, + config=config, + quant_config=MagicMock(), + ) + + qkv_parallel_linear_cls.assert_called_once() + projection_args = qkv_parallel_linear_cls.call_args.kwargs + self.assertEqual(projection_args["tp_rank"], 2) + self.assertEqual(projection_args["tp_size"], 4) + + self.assertEqual(column_parallel_linear_cls.call_count, 3) + for call in column_parallel_linear_cls.call_args_list: + self.assertEqual(call.kwargs["tp_rank"], 2) + self.assertEqual(call.kwargs["tp_size"], 4) + + merged_column_parallel_linear_cls.assert_called_once() + conv_args = merged_column_parallel_linear_cls.call_args.kwargs + self.assertEqual(conv_args["tp_rank"], 2) + self.assertEqual(conv_args["tp_size"], 4) + + row_parallel_linear_cls.assert_called_once() + output_args = row_parallel_linear_cls.call_args.kwargs + self.assertEqual(output_args["tp_rank"], 2) + self.assertEqual(output_args["tp_size"], 4) + self.assertFalse(output_args["reduce_results"]) + + def test_kda_disables_global_tp_fusion_when_cp_owns_fewer_heads(self): + config = SimpleNamespace( + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + }, + v_head_dim=128, + dtype=torch.bfloat16, + ) + + with ( + get_parallel().override( + tp_size=8, + tp_rank=6, + attn_tp_size=2, + attn_tp_rank=1, + attn_cp_size=4, + attn_cp_rank=2, + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelRepeatedLinear", + return_value=MagicMock(), + ) as fused_projection_cls, + patch( + "sglang.srt.models.kimi_linear.QKVParallelLinear", + return_value=MagicMock(), + ) as qkv_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.ReplicatedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.ColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.FusedRMSNormGated", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RowParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RadixLinearAttention", + return_value=MagicMock(), + ), + ): + attention = KimiDeltaAttention( + layer_idx=0, + hidden_size=256, + config=config, + ) + + self.assertFalse(attention.do_fuse_qkvbfg) + fused_projection_cls.assert_not_called() + qkv_parallel_linear_cls.assert_called_once() + + def test_kda_keeps_fusion_when_cp_matches_global_tp(self): + config = SimpleNamespace( + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + }, + v_head_dim=128, + dtype=torch.bfloat16, + ) with ( get_parallel().override( @@ -507,10 +657,68 @@ def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): tp_rank=2, attn_tp_size=1, attn_tp_rank=0, + attn_cp_size=4, + attn_cp_rank=2, + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelRepeatedLinear", + return_value=MagicMock(), + ) as fused_projection_cls, + patch( + "sglang.srt.models.kimi_linear.ColumnParallelBatchedLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.MergedColumnParallelLinear", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.FusedRMSNormGated", + return_value=MagicMock(), + ), + patch( + "sglang.srt.models.kimi_linear.RowParallelLinear", + return_value=MagicMock(), + ) as row_parallel_linear_cls, + patch( + "sglang.srt.models.kimi_linear.RadixLinearAttention", + return_value=MagicMock(), + ), + ): + attention = KimiDeltaAttention( + layer_idx=0, + hidden_size=256, + config=config, + ) + + self.assertTrue(attention.do_fuse_qkvbfg) + fused_projection_cls.assert_called_once() + self.assertEqual(attention.split_sizes, [3072, 8, 256]) + self.assertFalse(row_parallel_linear_cls.call_args.kwargs["reduce_results"]) + + def test_kda_keeps_global_tp_projection_ownership_without_cp(self): + config = SimpleNamespace( + linear_attn_config={ + "head_dim": 128, + "num_heads": 32, + "short_conv_kernel_size": 4, + }, + v_head_dim=128, + dtype=torch.bfloat16, + ) + + with ( + get_parallel().override( + tp_size=4, + tp_rank=2, + attn_tp_size=4, + attn_tp_rank=2, + attn_cp_size=1, + attn_cp_rank=0, ), patch( "sglang.srt.models.kimi_linear.QKVParallelLinear", - return_value=qkv_parallel_linear, + return_value=MagicMock(), ) as qkv_parallel_linear_cls, patch( "sglang.srt.models.kimi_linear.ReplicatedLinear", @@ -531,25 +739,27 @@ def test_unfused_kda_projection_uses_global_tp_rank_and_size(self): patch( "sglang.srt.models.kimi_linear.RowParallelLinear", return_value=MagicMock(), - ), + ) as row_parallel_linear_cls, patch( "sglang.srt.models.kimi_linear.RadixLinearAttention", return_value=MagicMock(), ), ): - KimiDeltaAttention( + attention = KimiDeltaAttention( layer_idx=0, hidden_size=256, config=config, quant_config=MagicMock(), ) - qkv_parallel_linear_cls.assert_called_once() projection_args = qkv_parallel_linear_cls.call_args.kwargs self.assertEqual(projection_args["tp_rank"], 2) self.assertEqual(projection_args["tp_size"], 4) + output_args = row_parallel_linear_cls.call_args.kwargs + self.assertTrue(output_args["reduce_results"]) + self.assertFalse(attention.kda_heads_use_cp) - def test_kda_cache_shape_uses_global_tp_size(self): + def test_kda_cache_shape_uses_cp_size(self): config = KimiLinearConfig( num_hidden_layers=2, linear_attn_config={ @@ -561,28 +771,92 @@ def test_kda_cache_shape_uses_global_tp_size(self): }, ) - with get_parallel().override(tp_size=4, attn_tp_size=1): + with get_parallel().override( + tp_size=8, + attn_tp_size=2, + attn_cp_size=4, + ): shape = config.mamba2_cache_params.shape self.assertEqual(shape.temporal, (8, 128, 128)) self.assertEqual(shape.conv, [(3, 3072)]) self.assertEqual(shape.num_k_heads_per_tp, 8) - def test_kda_state_weight_loader_uses_global_tp_rank(self): + def test_kda_state_weight_loader_uses_cp_rank(self): + param = torch.nn.Parameter(torch.zeros(2)) + loaded_weight = torch.arange(16, dtype=param.dtype) + loader = _kda_head_sharded_weight_loader(0) + + with get_parallel().override( + tp_size=8, + tp_rank=6, + attn_tp_size=2, + attn_tp_rank=1, + attn_cp_size=4, + attn_cp_rank=2, + ): + loader(param, loaded_weight) + + torch.testing.assert_close(param, torch.tensor([4.0, 5.0])) + + def test_kda_state_weight_loader_keeps_global_tp_without_cp(self): param = torch.nn.Parameter(torch.zeros(2)) loaded_weight = torch.arange(8, dtype=param.dtype) - loader = _global_tp_sharded_weight_loader(0) + loader = _kda_head_sharded_weight_loader(0) with get_parallel().override( tp_size=4, tp_rank=2, - attn_tp_size=1, - attn_tp_rank=0, + attn_tp_size=4, + attn_tp_rank=2, + attn_cp_size=1, + attn_cp_rank=0, ): loader(param, loaded_weight) torch.testing.assert_close(param, torch.tensor([4.0, 5.0])) + def test_kda_output_reduces_only_over_its_cp_head_group(self): + for kda_heads_use_cp in (True, False): + with self.subTest(kda_heads_use_cp=kda_heads_use_cp): + attention = object.__new__(KimiDeltaAttention) + torch.nn.Module.__init__(attention) + attention.do_fuse_qkvbfg = True + attention.head_dim = 2 + attention.kda_heads_use_cp = kda_heads_use_cp + attention.forward_qkvbfg_fused = MagicMock( + return_value=( + torch.zeros(2, 12), + torch.zeros(2, 2), + torch.zeros(2, 4), + torch.zeros(2, 4), + ) + ) + attention.attn = MagicMock(return_value=torch.zeros(1, 2, 2, 2)) + attention.o_norm = MagicMock(return_value=torch.zeros(1, 2, 2, 2)) + partial_output = torch.arange(8, dtype=torch.float32).view(2, 4) + attention.o_proj = MagicMock(return_value=(partial_output, None)) + cp_group = MagicMock() + cp_group.all_reduce.return_value = partial_output + 10 + forward_batch = SimpleNamespace( + forward_mode=SimpleNamespace(is_decode=lambda: True) + ) + + with get_parallel().override(attn_cp_group=cp_group): + output = attention( + hidden_states=torch.zeros(2, 4), + positions=torch.arange(2), + forward_batch=forward_batch, + zero_allocator=MagicMock(), + ) + + if kda_heads_use_cp: + cp_group.all_reduce.assert_called_once_with(partial_output) + torch.testing.assert_close(output, partial_output + 10) + else: + cp_group.all_reduce.assert_not_called() + torch.testing.assert_close(output, partial_output) + class TestKimiLinearCPV2Activation(CustomTestCase): def test_kimi_linear_uses_cp_v2_by_default(self): @@ -700,8 +974,8 @@ def run_attention(q_fused, forward_batch, device, attn_fn, **kwargs): head_dim=6, ) forward_batch = SimpleNamespace() - backend._run_cp_paged_attention.side_effect = ( - lambda wrapper, q_chunk, layer: q_chunk[..., : layer.v_head_dim] + backend._run_cp_paged_attention.side_effect = lambda wrapper, q_chunk, layer: ( + q_chunk[..., : layer.v_head_dim] ) with ( From 44e30b2609ec38c2bf7f10d9d985add6834ddea4 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sun, 19 Jul 2026 03:55:33 -0700 Subject: [PATCH 34/35] test: avoid CUDA stream in Kimi CP CPU test --- test/registered/cp/test_kimi_linear_cp_v2.py | 1 + 1 file changed, 1 insertion(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index cf15f16122f6..d8814ef27431 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -326,6 +326,7 @@ def _pad_rank_tensors(rank_tensors): local_hidden_states[rank], local_residuals[rank], forward_batches[rank], + stream=MagicMock(), ) torch.testing.assert_close(gathered_hidden, hidden_states) From fd7edf561e771cc09f53333b41d5de61fbc27b4c Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Sun, 19 Jul 2026 04:01:30 -0700 Subject: [PATCH 35/35] test: mock optional FlashInfer MLA wrapper --- test/registered/cp/test_kimi_linear_cp_v2.py | 1 + 1 file changed, 1 insertion(+) diff --git a/test/registered/cp/test_kimi_linear_cp_v2.py b/test/registered/cp/test_kimi_linear_cp_v2.py index d8814ef27431..4e46aa194659 100644 --- a/test/registered/cp/test_kimi_linear_cp_v2.py +++ b/test/registered/cp/test_kimi_linear_cp_v2.py @@ -930,6 +930,7 @@ def test_cp_wrapper_plan_uses_physical_token_page_size(self): patch( "sglang.srt.layers.attention.flashinfer_mla_backend.BatchMLAPagedAttentionWrapper", return_value=wrapper, + create=True, ), patch( "sglang.srt.layers.attention.flashinfer_mla_backend.create_flashinfer_kv_indices_triton",