Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
35 commits
Select commit Hold shift + click to select a range
340cc02
docs: design Kimi-Linear CP-v2 transitions
Fridge003 Jul 18, 2026
f3f6bf4
docs: plan Kimi-Linear CP-v2 implementation
Fridge003 Jul 18, 2026
1d54fc2
test: specify Kimi-Linear CP-v2 entry gather
Fridge003 Jul 18, 2026
407354f
feat: gather Kimi-Linear CP input for KDA
Fridge003 Jul 18, 2026
3f452cd
test: specify KDA to MLA CP split
Fridge003 Jul 18, 2026
630f45d
feat: shard Kimi-Linear state before MLA
Fridge003 Jul 18, 2026
d5921b9
test: specify MLA to KDA CP gather
Fridge003 Jul 18, 2026
3c213d6
feat: gather Kimi-Linear residual before KDA
Fridge003 Jul 18, 2026
7b24010
test: specify Kimi decoder CP communicator wiring
Fridge003 Jul 18, 2026
2d7e74d
feat: wire CP-v2 transitions into Kimi layers
Fridge003 Jul 18, 2026
7f9362e
test: complete Kimi decoder fixture
Fridge003 Jul 18, 2026
92ce1af
test: isolate Kimi MLA layer wiring
Fridge003 Jul 18, 2026
70a8767
test: require CP-v2 for Kimi-Linear
Fridge003 Jul 18, 2026
898ac51
feat: enable CP-v2 for Kimi-Linear
Fridge003 Jul 18, 2026
cd71780
test: require Kimi input embedding accessor
Fridge003 Jul 18, 2026
8754d7e
feat: expose Kimi input embeddings to CP-v2
Fridge003 Jul 18, 2026
78c2083
test: cover Kimi CP-v2 no-op and zigzag round trip
Fridge003 Jul 18, 2026
02f464b
test: require KDA heads to use global TP
Fridge003 Jul 18, 2026
4f6220a
fix: partition KDA backend heads over global TP
Fridge003 Jul 18, 2026
bf733d9
test: require unfused KDA projection to use global TP
Fridge003 Jul 18, 2026
f0364b5
fix: shard unfused KDA projections over global TP
Fridge003 Jul 18, 2026
01b92d0
test: require KDA cache state to use global TP
Fridge003 Jul 18, 2026
e2ecec6
fix: shard KDA state cache over global TP
Fridge003 Jul 18, 2026
55f2a64
test: cover Kimi CP-v2 embedding keyword
Fridge003 Jul 18, 2026
b41d63a
fix: accept CP-v2 input embeddings in Kimi
Fridge003 Jul 18, 2026
5af9dd8
test: cover FlashInfer MLA CP-v2 dispatch
Fridge003 Jul 18, 2026
34ced26
feat: support FlashInfer MLA in CP-v2
Fridge003 Jul 18, 2026
5ff97cc
test: match FlashInfer MLA output layout
Fridge003 Jul 18, 2026
37c791f
fix: run Kimi KDA and MLP on global TP batches
Fridge003 Jul 18, 2026
697fff4
docs: target merged CP-v2 base
Fridge003 Jul 18, 2026
c422342
fix: address Kimi CP-v2 review feedback
Fridge003 Jul 18, 2026
8e9d306
docs: remove implementation planning files
Fridge003 Jul 18, 2026
c0f5526
feat: shard KDA heads over CP ranks
Fridge003 Jul 19, 2026
44e30b2
test: avoid CUDA stream in Kimi CP CPU test
Fridge003 Jul 19, 2026
fd7edf5
test: mock optional FlashInfer MLA wrapper
Fridge003 Jul 19, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 18 additions & 2 deletions python/sglang/srt/configs/kimi_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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().attn_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"],
Expand Down
164 changes: 164 additions & 0 deletions python/sglang/srt/layers/attention/flashinfer_mla_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -377,6 +387,12 @@ 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,
Expand Down Expand Up @@ -511,6 +527,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,
1,
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,
Expand All @@ -522,6 +674,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
Expand Down
112 changes: 112 additions & 0 deletions python/sglang/srt/layers/cp/kimi_linear.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
# 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

import torch

from sglang.srt.layers.cp.utils import get_cp_strategy, is_cp_v2_active

if TYPE_CHECKING:
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],
is_last_layer: bool = False,
) -> None:
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,
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):
return hidden_states, residual

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
)
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:
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
if stream is None:
stream = torch.cuda.current_stream()
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]]:
Comment on lines +96 to +102

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

medium

The stream parameter in postprocess_layer is unused. It is recommended to remove it from the method signature to keep the API clean and maintainable.

    def postprocess_layer(
        self,
        hidden_states: torch.Tensor,
        residual: Optional[torch.Tensor],
        forward_batch: ForwardBatch,
    ) -> 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
1 change: 1 addition & 0 deletions python/sglang/srt/layers/cp/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
{
"Qwen3MoeForCausalLM",
"DeepseekV3ForCausalLM",
"KimiLinearForCausalLM",
}
)

Expand Down
Loading
Loading