Support DCP for Kimi Linear model (#32612)
Co-authored-by: Julien Lin <jullin@nvidia.com> Co-authored-by: kpham-sgl <khoa.pham@radixark.ai>
This commit is contained in:
co-authored by
Julien Lin
kpham-sgl
parent
c4fc241fd3
commit
ef6c07008b
@@ -76,6 +76,40 @@ def create_triton_kv_indices_for_dcp_triton(
|
|||||||
# KV-index build (PR #14194, MLA): global prefix+extend layout for the
|
# KV-index build (PR #14194, MLA): global prefix+extend layout for the
|
||||||
# all-gathered dcp_kv_buffer, plus the per-rank shard/compact kernel.
|
# all-gathered dcp_kv_buffer, plus the per-rank shard/compact kernel.
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
@triton.jit
|
||||||
|
def create_mla_kv_page_table_for_dcp(
|
||||||
|
req_to_token_ptr,
|
||||||
|
req_pool_indices_ptr,
|
||||||
|
local_seq_lens_ptr,
|
||||||
|
block_kv_indices_ptr,
|
||||||
|
req_to_token_stride: tl.constexpr,
|
||||||
|
block_table_stride: tl.constexpr,
|
||||||
|
PHYSICAL_PAGE_SIZE: tl.constexpr,
|
||||||
|
DCP_SIZE: tl.constexpr,
|
||||||
|
DCP_RANK: tl.constexpr,
|
||||||
|
PAGES_PER_BLOCK: tl.constexpr,
|
||||||
|
):
|
||||||
|
req = tl.program_id(0)
|
||||||
|
page_block = tl.program_id(1)
|
||||||
|
page_offsets = page_block * PAGES_PER_BLOCK + tl.arange(0, PAGES_PER_BLOCK)
|
||||||
|
local_len = tl.load(local_seq_lens_ptr + req)
|
||||||
|
local_pages = tl.cdiv(local_len, PHYSICAL_PAGE_SIZE)
|
||||||
|
mask = page_offsets < local_pages
|
||||||
|
global_positions = DCP_RANK + page_offsets * PHYSICAL_PAGE_SIZE * DCP_SIZE
|
||||||
|
req_pool_index = tl.load(req_pool_indices_ptr + req)
|
||||||
|
virtual_locs = tl.load(
|
||||||
|
req_to_token_ptr + req_pool_index * req_to_token_stride + global_positions,
|
||||||
|
mask=mask,
|
||||||
|
other=0,
|
||||||
|
)
|
||||||
|
physical_pages = virtual_locs // DCP_SIZE // PHYSICAL_PAGE_SIZE
|
||||||
|
tl.store(
|
||||||
|
block_kv_indices_ptr + req * block_table_stride + page_offsets,
|
||||||
|
physical_pages,
|
||||||
|
mask=mask,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@triton.jit
|
@triton.jit
|
||||||
def create_dcp_kv_indices(
|
def create_dcp_kv_indices(
|
||||||
kv_indptr,
|
kv_indptr,
|
||||||
|
|||||||
@@ -89,9 +89,9 @@ def create_tokenspeed_mla_backend(runner):
|
|||||||
def create_cutedsl_mla_backend(runner):
|
def create_cutedsl_mla_backend(runner):
|
||||||
if not runner.use_mla_backend:
|
if not runner.use_mla_backend:
|
||||||
raise ValueError("cutedsl_mla backend can only be used with MLA models.")
|
raise ValueError("cutedsl_mla backend can only be used with MLA models.")
|
||||||
from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend
|
from sglang.srt.layers.attention.cutedsl_mla_backend import CuteDslMLABackend
|
||||||
|
|
||||||
return TRTLLMMLABackend(runner, backend="cute-dsl")
|
return CuteDslMLABackend(runner)
|
||||||
|
|
||||||
|
|
||||||
@register_attention_backend("aiter")
|
@register_attention_backend("aiter")
|
||||||
|
|||||||
@@ -0,0 +1,451 @@
|
|||||||
|
"""
|
||||||
|
Attention backend for the flashinfer cute-dsl MLA decode kernels with decode
|
||||||
|
context parallelism (DCP).
|
||||||
|
|
||||||
|
Subclasses :class:`TRTLLMMLABackend` (``backend="cute-dsl"``) to reuse its MLA
|
||||||
|
data preparation, workspace, and prefill plumbing. The flashinfer cute-dsl
|
||||||
|
monolithic MLA decode kernel natively accepts cyclic DCP metadata
|
||||||
|
(``enable_dcp`` / ``cp_world`` / ``cp_rank`` / ``causal_seqlens_kv_global``) and
|
||||||
|
returns the rank-local ``(out, lse)`` needed by the cross-rank merge in
|
||||||
|
``deepseek_common/attention_forward_methods/forward_mla.py``.
|
||||||
|
|
||||||
|
Non-DCP (``dcp_size == 1``) decode falls through to the base cute-dsl path
|
||||||
|
unchanged. The DCP metadata helpers below are intentionally duplicated from
|
||||||
|
:mod:`tokenspeed_mla_backend` (they are kernel-agnostic) so that TokenSpeed
|
||||||
|
stays untouched; both should collapse into the base once the cute-dsl decode
|
||||||
|
path is stable (see the TODO in tokenspeed_mla_backend.py).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import math
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||||
|
create_mla_kv_page_table_for_dcp,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
|
||||||
|
from sglang.kernels.ops.attention.utils import (
|
||||||
|
concat_mla_absorb_q_general,
|
||||||
|
mla_quantize_and_rope_for_fp8,
|
||||||
|
mla_quantize_without_rope_for_fp8,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
|
get_num_page_per_block_flashmla,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||||
|
TRTLLMMLABackend,
|
||||||
|
TRTLLMMLAMultiStepDraftBackend,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.utils import is_flashinfer_available
|
||||||
|
|
||||||
|
if is_flashinfer_available():
|
||||||
|
import flashinfer
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# The flashinfer cute-dsl MLA decode kernel returns a natural-log (base-e) LSE,
|
||||||
|
# whereas sglang's DCP cross-rank merge (forward_mla: dcp_a2a_lse_reduce /
|
||||||
|
# cp_lse_ag_out_rs_mla) assumes the FlashInfer-MLA/FlashMLA base-2 convention
|
||||||
|
# (is_lse_base_on_e=False). Multiplying a natural-log LSE by log2(e) rebases it
|
||||||
|
# to base-2 (the softmax output is base-invariant; only the LSE value changes).
|
||||||
|
# CONFIRMED base-e (not base-2), so this rebase is required, not optional:
|
||||||
|
# the flashinfer-dcp-backport public-API unit test asserts the public
|
||||||
|
# trtllm_batch_decode_with_kv_cache_mla LSE against a torch.logsumexp
|
||||||
|
# (natural-log) reference at atol=1e-2 and passes (a base-2 LSE would be
|
||||||
|
# off by 1/ln2 ~= 44%). GPU job 467640:
|
||||||
|
# tests/attention/test_cute_dsl_mla_dcp*.py 27/27 + 17/17 pass.
|
||||||
|
_LSE_BASE2_FROM_NATURAL_LOG = math.log2(math.e)
|
||||||
|
|
||||||
|
|
||||||
|
class CuteDslMLABackend(TRTLLMMLABackend):
|
||||||
|
"""flashinfer cute-dsl MLA decode backend with decode context parallelism."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
model_runner: ModelRunner,
|
||||||
|
skip_prefill: bool = False,
|
||||||
|
kv_indptr_buf: Optional[torch.Tensor] = None,
|
||||||
|
q_indptr_decode_buf: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
model_runner,
|
||||||
|
skip_prefill,
|
||||||
|
kv_indptr_buf,
|
||||||
|
q_indptr_decode_buf,
|
||||||
|
backend="cute-dsl",
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# DCP metadata (rank-local KV lengths + page table).
|
||||||
|
# Duplicated from TokenspeedMLABackend — kernel-agnostic, keyed only on
|
||||||
|
# dcp_size / dcp_rank / page_size / req_to_token.
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return seq_lens
|
||||||
|
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
|
||||||
|
torch.int32
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return max_seq_len
|
||||||
|
local_max = max_seq_len // parallel.dcp_size + int(
|
||||||
|
parallel.dcp_rank < max_seq_len % parallel.dcp_size
|
||||||
|
)
|
||||||
|
# A positive scheduling bound is required even when every sequence in a
|
||||||
|
# padded graph row is empty on this rank.
|
||||||
|
return max(local_max, 1)
|
||||||
|
|
||||||
|
def _fill_dcp_block_kv_indices(
|
||||||
|
self,
|
||||||
|
block_kv_indices: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
local_seq_lens: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
parallel = get_parallel()
|
||||||
|
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
|
||||||
|
create_mla_kv_page_table_for_dcp[
|
||||||
|
(
|
||||||
|
block_kv_indices.shape[0],
|
||||||
|
get_num_kv_index_blocks_flashmla(
|
||||||
|
block_kv_indices.shape[1], self.page_size
|
||||||
|
),
|
||||||
|
)
|
||||||
|
](
|
||||||
|
self.req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
local_seq_lens,
|
||||||
|
block_kv_indices,
|
||||||
|
self.req_to_token.stride(0),
|
||||||
|
block_kv_indices.stride(0),
|
||||||
|
PHYSICAL_PAGE_SIZE=self.page_size,
|
||||||
|
DCP_SIZE=parallel.dcp_size,
|
||||||
|
DCP_RANK=parallel.dcp_rank,
|
||||||
|
PAGES_PER_BLOCK=pages_per_block,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _create_block_kv_indices(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
max_blocks: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
device: torch.device,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if not get_parallel().dcp_enabled:
|
||||||
|
return super()._create_block_kv_indices(
|
||||||
|
batch_size,
|
||||||
|
max_blocks,
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
device,
|
||||||
|
)
|
||||||
|
block_kv_indices = torch.full(
|
||||||
|
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
block_kv_indices,
|
||||||
|
req_pool_indices,
|
||||||
|
self._get_dcp_local_seq_lens(seq_lens),
|
||||||
|
)
|
||||||
|
return block_kv_indices
|
||||||
|
|
||||||
|
def _init_cuda_graph_metadata(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
num_tokens: int,
|
||||||
|
forward_mode,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
device: torch.device,
|
||||||
|
):
|
||||||
|
super()._init_cuda_graph_metadata(
|
||||||
|
bs, num_tokens, forward_mode, seq_lens, device
|
||||||
|
)
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
if forward_mode.is_target_verify():
|
||||||
|
self.forward_decode_metadata.global_seq_lens_k = torch.zeros_like(
|
||||||
|
self.forward_decode_metadata.seq_lens_k
|
||||||
|
)
|
||||||
|
self.forward_decode_metadata.max_seq_len_k = (
|
||||||
|
self._get_dcp_local_max_seq_len(
|
||||||
|
self.max_context_len
|
||||||
|
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _apply_cuda_graph_metadata(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
forward_mode,
|
||||||
|
):
|
||||||
|
if not get_parallel().dcp_enabled:
|
||||||
|
return super()._apply_cuda_graph_metadata(
|
||||||
|
bs,
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
forward_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
|
if forward_mode.is_target_verify():
|
||||||
|
torch.add(
|
||||||
|
seq_lens[:bs],
|
||||||
|
self.num_draft_tokens,
|
||||||
|
out=metadata.global_seq_lens_k,
|
||||||
|
)
|
||||||
|
metadata.seq_lens_k.copy_(
|
||||||
|
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
|
||||||
|
)
|
||||||
|
local_seq_lens = metadata.seq_lens_k
|
||||||
|
elif forward_mode.is_draft_extend_v2():
|
||||||
|
num_tokens_per_req = self.num_draft_tokens
|
||||||
|
metadata.max_seq_len_q = num_tokens_per_req
|
||||||
|
metadata.sum_seq_lens_q = num_tokens_per_req * bs
|
||||||
|
seq_lens = seq_lens[:bs]
|
||||||
|
metadata.seq_lens_k.copy_(seq_lens)
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||||
|
else:
|
||||||
|
seq_lens = seq_lens[:bs]
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||||
|
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
metadata.block_kv_indices,
|
||||||
|
req_pool_indices[:bs],
|
||||||
|
local_seq_lens,
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
super().init_forward_metadata(forward_batch)
|
||||||
|
if (
|
||||||
|
get_parallel().dcp_enabled
|
||||||
|
and self.forward_decode_metadata is not None
|
||||||
|
and (
|
||||||
|
forward_batch.forward_mode.is_decode_or_idle()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
)
|
||||||
|
):
|
||||||
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
metadata = self.forward_decode_metadata
|
||||||
|
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||||
|
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||||
|
metadata.global_seq_lens_k
|
||||||
|
)
|
||||||
|
self.forward_decode_metadata.max_seq_len_k = (
|
||||||
|
self._get_dcp_local_max_seq_len(
|
||||||
|
self.forward_decode_metadata.max_seq_len_k
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Kernel + decode forward.
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def _run_decode_kernel(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
kv_cache: torch.Tensor,
|
||||||
|
block_tables: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
max_seq_len: int,
|
||||||
|
layer: RadixAttention,
|
||||||
|
*,
|
||||||
|
causal_seqs: Optional[torch.Tensor] = None,
|
||||||
|
cp_world: int = 1,
|
||||||
|
cp_rank: int = 0,
|
||||||
|
return_lse: bool = False,
|
||||||
|
):
|
||||||
|
"""Call the flashinfer cute-dsl MLA decode kernel.
|
||||||
|
|
||||||
|
Without DCP (``cp_world <= 1``) this defers to the base cute-dsl path.
|
||||||
|
With DCP, ``seq_lens`` are this rank's cyclic-local KV lengths and
|
||||||
|
``causal_seqs`` the global per-request KV lengths; the kernel returns a
|
||||||
|
rank-local ``(out, lse)`` (LSE rebased to base-2 for the sglang merge).
|
||||||
|
"""
|
||||||
|
if cp_world <= 1:
|
||||||
|
return super()._run_decode_kernel(
|
||||||
|
query, kv_cache, block_tables, seq_lens, max_seq_len, layer
|
||||||
|
)
|
||||||
|
if causal_seqs is None:
|
||||||
|
raise ValueError(
|
||||||
|
"causal_seqs (global per-request KV lengths) is required for DCP "
|
||||||
|
"MLA decode."
|
||||||
|
)
|
||||||
|
bmm1_scale = self._compute_decode_bmm1_scale(layer)
|
||||||
|
raw_out, lse = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla(
|
||||||
|
query=query,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
workspace_buffer=self.workspace_buffer,
|
||||||
|
qk_nope_head_dim=self.qk_nope_head_dim,
|
||||||
|
kv_lora_rank=self.kv_lora_rank,
|
||||||
|
qk_rope_head_dim=self.qk_rope_head_dim,
|
||||||
|
block_tables=block_tables,
|
||||||
|
seq_lens=(
|
||||||
|
seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32)
|
||||||
|
),
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
bmm1_scale=bmm1_scale,
|
||||||
|
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
||||||
|
backend="cute-dsl",
|
||||||
|
enable_dcp=True,
|
||||||
|
cp_world=cp_world,
|
||||||
|
cp_rank=cp_rank,
|
||||||
|
causal_seqlens_kv_global=(
|
||||||
|
causal_seqs
|
||||||
|
if causal_seqs.dtype == torch.int32
|
||||||
|
else causal_seqs.to(torch.int32)
|
||||||
|
),
|
||||||
|
return_lse=True, # DCP requires the rank-local LSE for the merge
|
||||||
|
)
|
||||||
|
return raw_out, lse * _LSE_BASE2_FROM_NATURAL_LOG
|
||||||
|
|
||||||
|
def forward_decode(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor, # q_nope
|
||||||
|
k: torch.Tensor, # k_nope
|
||||||
|
v: torch.Tensor, # not used in this backend
|
||||||
|
layer: RadixAttention,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
save_kv_cache: bool = True,
|
||||||
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
|
cos_sin_cache: Optional[torch.Tensor] = None,
|
||||||
|
is_neox: Optional[bool] = False,
|
||||||
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return super().forward_decode(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
layer,
|
||||||
|
forward_batch,
|
||||||
|
save_kv_cache,
|
||||||
|
q_rope,
|
||||||
|
k_rope,
|
||||||
|
cos_sin_cache,
|
||||||
|
is_neox,
|
||||||
|
llama_4_scaling,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Query / KV preparation mirrors the base cute-dsl decode (both FP16 and
|
||||||
|
# FP8 KV), then swaps to the DCP kernel call + rank-local return.
|
||||||
|
merge_query = q_rope is not None
|
||||||
|
if self.data_type == torch.float8_e4m3fn:
|
||||||
|
assert q_rope is not None and k_rope is not None
|
||||||
|
if cos_sin_cache is None:
|
||||||
|
q, k, k_rope = mla_quantize_without_rope_for_fp8(
|
||||||
|
q, q_rope, k.squeeze(1), k_rope.squeeze(1)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
q, k, k_rope = mla_quantize_and_rope_for_fp8(
|
||||||
|
q,
|
||||||
|
q_rope,
|
||||||
|
k.squeeze(1),
|
||||||
|
k_rope.squeeze(1),
|
||||||
|
forward_batch.positions,
|
||||||
|
cos_sin_cache,
|
||||||
|
is_neox,
|
||||||
|
self.kv_lora_rank,
|
||||||
|
self.qk_rope_head_dim,
|
||||||
|
)
|
||||||
|
merge_query = False
|
||||||
|
|
||||||
|
if save_kv_cache:
|
||||||
|
assert k is not None and k_rope is not None
|
||||||
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
|
layer, forward_batch.out_cache_loc, k, k_rope
|
||||||
|
)
|
||||||
|
|
||||||
|
if merge_query:
|
||||||
|
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||||
|
q_rope_reshaped = q_rope.view(
|
||||||
|
-1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim
|
||||||
|
)
|
||||||
|
query = concat_mla_absorb_q_general(q_nope, q_rope_reshaped)
|
||||||
|
else:
|
||||||
|
query = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||||
|
|
||||||
|
if llama_4_scaling is not None:
|
||||||
|
query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type)
|
||||||
|
if query.dim() == 3:
|
||||||
|
query = query.unsqueeze(1)
|
||||||
|
|
||||||
|
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
|
||||||
|
|
||||||
|
metadata = (
|
||||||
|
getattr(forward_batch, "decode_trtllm_mla_metadata", None)
|
||||||
|
or self.forward_decode_metadata
|
||||||
|
)
|
||||||
|
metadata_batch_size = getattr(metadata, "batch_size", None)
|
||||||
|
if (
|
||||||
|
metadata_batch_size is not None
|
||||||
|
and metadata_batch_size < forward_batch.batch_size
|
||||||
|
):
|
||||||
|
self.init_forward_metadata(forward_batch)
|
||||||
|
metadata = forward_batch.decode_trtllm_mla_metadata
|
||||||
|
|
||||||
|
global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size]
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens)
|
||||||
|
raw_out, lse = self._run_decode_kernel(
|
||||||
|
query=query,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
block_tables=metadata.block_kv_indices,
|
||||||
|
seq_lens=local_seq_lens,
|
||||||
|
max_seq_len=metadata.max_seq_len_k,
|
||||||
|
layer=layer,
|
||||||
|
causal_seqs=global_seq_lens,
|
||||||
|
cp_world=parallel.dcp_size,
|
||||||
|
cp_rank=parallel.dcp_rank,
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
output = raw_out.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||||
|
lse = lse.view(-1, layer.tp_q_head_num)
|
||||||
|
# Zero-KV rows (a request this rank owns no cyclic slice for) get a
|
||||||
|
# neutral (out=0, lse=-inf) state so the cross-rank merge ignores them.
|
||||||
|
fixup_zero_kv_rows(
|
||||||
|
output,
|
||||||
|
lse,
|
||||||
|
local_seq_lens,
|
||||||
|
self.q_indptr_decode[: forward_batch.batch_size + 1],
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
return output.flatten(1), lse
|
||||||
|
|
||||||
|
|
||||||
|
class CuteDslMLAMultiStepDraftBackend(TRTLLMMLAMultiStepDraftBackend):
|
||||||
|
"""Multi-step draft backend for cutedsl_mla used by EAGLE / DSPARK."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
|
||||||
|
):
|
||||||
|
super().__init__(model_runner, topk, speculative_num_steps)
|
||||||
|
# Parent populates self.attn_backends with TRT-LLM instances; replace
|
||||||
|
# them with cute-dsl instances sharing the parent's index buffers.
|
||||||
|
for i in range(self.speculative_num_steps - 1):
|
||||||
|
self.attn_backends[i] = CuteDslMLABackend(
|
||||||
|
model_runner,
|
||||||
|
skip_prefill=True,
|
||||||
|
kv_indptr_buf=self.kv_indptr[i],
|
||||||
|
q_indptr_decode_buf=self.q_indptr_decode,
|
||||||
|
)
|
||||||
@@ -49,11 +49,11 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
generate_draft_decode_kv_indices,
|
generate_draft_decode_kv_indices,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
|
get_cuda_graph_max_batch_size,
|
||||||
get_int_env_var,
|
get_int_env_var,
|
||||||
is_flashinfer_available,
|
is_flashinfer_available,
|
||||||
is_sm100_supported,
|
is_sm100_supported,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
require_gathered_buffer,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -63,18 +63,6 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _cuda_graph_capture_max_bs(server_args, max_bs: int) -> int:
|
|
||||||
"""Pad max_bs to the alignment cuda-graph capture uses (see get_batch_sizes_to_capture)."""
|
|
||||||
mul_base = 1
|
|
||||||
if server_args.enable_two_batch_overlap:
|
|
||||||
mul_base *= 2
|
|
||||||
if require_gathered_buffer(server_args):
|
|
||||||
mul_base *= get_parallel().attn_tp_size
|
|
||||||
if mul_base % get_parallel().attn_cp_size != 0:
|
|
||||||
mul_base *= get_parallel().attn_cp_size
|
|
||||||
return (max_bs + mul_base - 1) // mul_base * mul_base
|
|
||||||
|
|
||||||
|
|
||||||
if envs.SGLANG_ENABLE_TORCH_COMPILE.get():
|
if envs.SGLANG_ENABLE_TORCH_COMPILE.get():
|
||||||
torch._logging.set_logs(dynamo=logging.ERROR)
|
torch._logging.set_logs(dynamo=logging.ERROR)
|
||||||
torch._dynamo.config.suppress_errors = True
|
torch._dynamo.config.suppress_errors = True
|
||||||
@@ -447,7 +435,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.workspace_buffer = global_workspace_buffer
|
self.workspace_buffer = global_workspace_buffer
|
||||||
max_bs = _cuda_graph_capture_max_bs(
|
max_bs = get_cuda_graph_max_batch_size(
|
||||||
model_runner.server_args, model_runner.req_to_token_pool.size
|
model_runner.server_args, model_runner.req_to_token_pool.size
|
||||||
)
|
)
|
||||||
if kv_indptr_buf is None:
|
if kv_indptr_buf is None:
|
||||||
@@ -2230,7 +2218,7 @@ class FlashInferMultiStepDraftBackend:
|
|||||||
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
|
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
|
||||||
max_bs = _cuda_graph_capture_max_bs(
|
max_bs = get_cuda_graph_max_batch_size(
|
||||||
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
|
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
|
||||||
)
|
)
|
||||||
self.kv_indptr = torch.zeros(
|
self.kv_indptr = torch.zeros(
|
||||||
|
|||||||
@@ -22,10 +22,10 @@ from __future__ import annotations
|
|||||||
|
|
||||||
"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell.
|
"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell.
|
||||||
|
|
||||||
Subclasses :class:`TRTLLMMLABackend` and overrides only ``_run_decode_kernel``
|
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation and
|
||||||
and ``_run_prefill_kernel``. All metadata, KV-cache layout, CUDA-graph
|
prefill plumbing. Decode-context parallelism is implemented here because the
|
||||||
plumbing, FP8 quantize/rope, draft-extend padding, and chunked-prefix
|
TokenSpeed decode kernel natively accepts CP rank/world metadata and returns
|
||||||
dispatch are inherited unchanged from the parent.
|
the partial log-sum-exp needed by the cross-rank merge.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -34,14 +34,28 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.kernels.jit.utils import is_arch_support_pdl
|
from sglang.kernels.jit.utils import is_arch_support_pdl
|
||||||
|
from sglang.kernels.ops.attention.dcp_kernels import (
|
||||||
|
create_mla_kv_page_table_for_dcp,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
|
||||||
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
|
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
|
||||||
mla_kv_pack_quantize_fp8,
|
mla_kv_pack_quantize_fp8,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.attention.utils import (
|
||||||
|
mla_quantize_and_rope_for_fp8,
|
||||||
|
mla_quantize_without_rope_for_fp8,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.kvcache.kv_indices import (
|
||||||
|
get_num_kv_index_blocks_flashmla,
|
||||||
|
get_num_page_per_block_flashmla,
|
||||||
|
)
|
||||||
from sglang.kernels.ops.quantization.fp8_quantize import fp8_quantize
|
from sglang.kernels.ops.quantization.fp8_quantize import fp8_quantize
|
||||||
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||||
TRTLLMMLABackend,
|
TRTLLMMLABackend,
|
||||||
TRTLLMMLAMultiStepDraftBackend,
|
TRTLLMMLAMultiStepDraftBackend,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import is_flashinfer_available, is_tokenspeed_mla_available
|
from sglang.srt.utils import is_flashinfer_available, is_tokenspeed_mla_available
|
||||||
|
|
||||||
if is_flashinfer_available():
|
if is_flashinfer_available():
|
||||||
@@ -273,6 +287,162 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
k_nope, k_pe, v, enable_pdl=is_arch_support_pdl()
|
k_nope, k_pe, v, enable_pdl=is_arch_support_pdl()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return seq_lens
|
||||||
|
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
|
||||||
|
torch.int32
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return max_seq_len
|
||||||
|
local_max = max_seq_len // parallel.dcp_size + int(
|
||||||
|
parallel.dcp_rank < max_seq_len % parallel.dcp_size
|
||||||
|
)
|
||||||
|
# TokenSpeed requires a positive scheduling bound even when every
|
||||||
|
# sequence in a padded graph row is empty on this rank.
|
||||||
|
return max(local_max, 1)
|
||||||
|
|
||||||
|
def _fill_dcp_block_kv_indices(
|
||||||
|
self,
|
||||||
|
block_kv_indices: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
local_seq_lens: torch.Tensor,
|
||||||
|
) -> None:
|
||||||
|
parallel = get_parallel()
|
||||||
|
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
|
||||||
|
create_mla_kv_page_table_for_dcp[
|
||||||
|
(
|
||||||
|
block_kv_indices.shape[0],
|
||||||
|
get_num_kv_index_blocks_flashmla(
|
||||||
|
block_kv_indices.shape[1], self.page_size
|
||||||
|
),
|
||||||
|
)
|
||||||
|
](
|
||||||
|
self.req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
local_seq_lens,
|
||||||
|
block_kv_indices,
|
||||||
|
self.req_to_token.stride(0),
|
||||||
|
block_kv_indices.stride(0),
|
||||||
|
PHYSICAL_PAGE_SIZE=self.page_size,
|
||||||
|
DCP_SIZE=parallel.dcp_size,
|
||||||
|
DCP_RANK=parallel.dcp_rank,
|
||||||
|
PAGES_PER_BLOCK=pages_per_block,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _create_block_kv_indices(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
max_blocks: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
device: torch.device,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if not get_parallel().dcp_enabled:
|
||||||
|
return super()._create_block_kv_indices(
|
||||||
|
batch_size,
|
||||||
|
max_blocks,
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
device,
|
||||||
|
)
|
||||||
|
|
||||||
|
block_kv_indices = torch.full(
|
||||||
|
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
block_kv_indices,
|
||||||
|
req_pool_indices,
|
||||||
|
self._get_dcp_local_seq_lens(seq_lens),
|
||||||
|
)
|
||||||
|
return block_kv_indices
|
||||||
|
|
||||||
|
def _init_cuda_graph_metadata(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
num_tokens: int,
|
||||||
|
forward_mode,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
device: torch.device,
|
||||||
|
):
|
||||||
|
super()._init_cuda_graph_metadata(
|
||||||
|
bs, num_tokens, forward_mode, seq_lens, device
|
||||||
|
)
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
self.forward_decode_metadata.max_seq_len_k = (
|
||||||
|
self._get_dcp_local_max_seq_len(
|
||||||
|
self.max_context_len
|
||||||
|
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _apply_cuda_graph_metadata(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
forward_mode,
|
||||||
|
):
|
||||||
|
if not get_parallel().dcp_enabled:
|
||||||
|
return super()._apply_cuda_graph_metadata(
|
||||||
|
bs, req_pool_indices, seq_lens, forward_mode
|
||||||
|
)
|
||||||
|
|
||||||
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
|
if forward_mode.is_target_verify():
|
||||||
|
torch.add(
|
||||||
|
seq_lens[:bs],
|
||||||
|
self.num_draft_tokens,
|
||||||
|
out=metadata.global_seq_lens_k,
|
||||||
|
)
|
||||||
|
metadata.seq_lens_k.copy_(
|
||||||
|
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
|
||||||
|
)
|
||||||
|
local_seq_lens = metadata.seq_lens_k
|
||||||
|
elif forward_mode.is_draft_extend_v2():
|
||||||
|
num_tokens_per_req = self.num_draft_tokens
|
||||||
|
metadata.max_seq_len_q = num_tokens_per_req
|
||||||
|
metadata.sum_seq_lens_q = num_tokens_per_req * bs
|
||||||
|
seq_lens = seq_lens[:bs]
|
||||||
|
metadata.seq_lens_k.copy_(seq_lens)
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||||
|
else:
|
||||||
|
seq_lens = seq_lens[:bs]
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
|
||||||
|
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
metadata.block_kv_indices,
|
||||||
|
req_pool_indices[:bs],
|
||||||
|
local_seq_lens,
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
|
super().init_forward_metadata(forward_batch)
|
||||||
|
if (
|
||||||
|
get_parallel().dcp_enabled
|
||||||
|
and self.forward_decode_metadata is not None
|
||||||
|
and (
|
||||||
|
forward_batch.forward_mode.is_decode_or_idle()
|
||||||
|
or forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
)
|
||||||
|
):
|
||||||
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
|
metadata = self.forward_decode_metadata
|
||||||
|
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||||
|
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||||
|
metadata.global_seq_lens_k
|
||||||
|
)
|
||||||
|
self.forward_decode_metadata.max_seq_len_k = (
|
||||||
|
self._get_dcp_local_max_seq_len(
|
||||||
|
self.forward_decode_metadata.max_seq_len_k
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def _run_decode_kernel(
|
def _run_decode_kernel(
|
||||||
self,
|
self,
|
||||||
query: torch.Tensor,
|
query: torch.Tensor,
|
||||||
@@ -281,7 +451,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
max_seq_len: int,
|
max_seq_len: int,
|
||||||
layer: RadixAttention,
|
layer: RadixAttention,
|
||||||
) -> torch.Tensor:
|
*,
|
||||||
|
causal_seqs: Optional[torch.Tensor] = None,
|
||||||
|
cp_world: int = 1,
|
||||||
|
cp_rank: int = 0,
|
||||||
|
return_lse: bool = False,
|
||||||
|
):
|
||||||
k_scale = getattr(layer, "k_scale_float", None)
|
k_scale = getattr(layer, "k_scale_float", None)
|
||||||
if k_scale is None:
|
if k_scale is None:
|
||||||
k_scale = 1.0
|
k_scale = 1.0
|
||||||
@@ -303,8 +478,111 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
softmax_scale=softmax_scale,
|
softmax_scale=softmax_scale,
|
||||||
output_scale=output_scale,
|
output_scale=output_scale,
|
||||||
enable_pdl=is_arch_support_pdl(),
|
enable_pdl=is_arch_support_pdl(),
|
||||||
|
return_lse=return_lse,
|
||||||
|
causal_seqs=causal_seqs,
|
||||||
|
cp_world=cp_world,
|
||||||
|
cp_rank=cp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def forward_decode(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
layer: RadixAttention,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
save_kv_cache: bool = True,
|
||||||
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
|
cos_sin_cache: Optional[torch.Tensor] = None,
|
||||||
|
is_neox: Optional[bool] = False,
|
||||||
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
|
):
|
||||||
|
parallel = get_parallel()
|
||||||
|
if not parallel.dcp_enabled:
|
||||||
|
return super().forward_decode(
|
||||||
|
q,
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
layer,
|
||||||
|
forward_batch,
|
||||||
|
save_kv_cache,
|
||||||
|
q_rope,
|
||||||
|
k_rope,
|
||||||
|
cos_sin_cache,
|
||||||
|
is_neox,
|
||||||
|
llama_4_scaling,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert q_rope is not None and k_rope is not None
|
||||||
|
if cos_sin_cache is None:
|
||||||
|
q, k, k_rope = mla_quantize_without_rope_for_fp8(
|
||||||
|
q, q_rope, k.squeeze(1), k_rope.squeeze(1)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
q, k, k_rope = mla_quantize_and_rope_for_fp8(
|
||||||
|
q,
|
||||||
|
q_rope,
|
||||||
|
k.squeeze(1),
|
||||||
|
k_rope.squeeze(1),
|
||||||
|
forward_batch.positions,
|
||||||
|
cos_sin_cache,
|
||||||
|
is_neox,
|
||||||
|
self.kv_lora_rank,
|
||||||
|
self.qk_rope_head_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
if save_kv_cache:
|
||||||
|
self.token_to_kv_pool.set_mla_kv_buffer(
|
||||||
|
layer, forward_batch.out_cache_loc, k, k_rope
|
||||||
|
)
|
||||||
|
|
||||||
|
query = q.view(-1, layer.tp_q_head_num, layer.head_dim)
|
||||||
|
if llama_4_scaling is not None:
|
||||||
|
query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type)
|
||||||
|
if query.dim() == 3:
|
||||||
|
query = query.unsqueeze(1)
|
||||||
|
|
||||||
|
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
|
||||||
|
metadata = (
|
||||||
|
getattr(forward_batch, "decode_trtllm_mla_metadata", None)
|
||||||
|
or self.forward_decode_metadata
|
||||||
|
)
|
||||||
|
metadata_batch_size = getattr(metadata, "batch_size", None)
|
||||||
|
if (
|
||||||
|
metadata_batch_size is not None
|
||||||
|
and metadata_batch_size < forward_batch.batch_size
|
||||||
|
):
|
||||||
|
self.init_forward_metadata(forward_batch)
|
||||||
|
metadata = forward_batch.decode_trtllm_mla_metadata
|
||||||
|
|
||||||
|
global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size]
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens)
|
||||||
|
raw_out, lse = self._run_decode_kernel(
|
||||||
|
query=query,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
block_tables=metadata.block_kv_indices,
|
||||||
|
seq_lens=local_seq_lens,
|
||||||
|
max_seq_len=metadata.max_seq_len_k,
|
||||||
|
layer=layer,
|
||||||
|
causal_seqs=global_seq_lens,
|
||||||
|
cp_world=parallel.dcp_size,
|
||||||
|
cp_rank=parallel.dcp_rank,
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
output = raw_out.view(-1, layer.tp_q_head_num, layer.v_head_dim)
|
||||||
|
lse = lse.view(-1, layer.tp_q_head_num)
|
||||||
|
fixup_zero_kv_rows(
|
||||||
|
output,
|
||||||
|
lse,
|
||||||
|
local_seq_lens,
|
||||||
|
self.q_indptr_decode[: forward_batch.batch_size + 1],
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
return output.flatten(1), lse
|
||||||
|
|
||||||
def _run_prefill_kernel(
|
def _run_prefill_kernel(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -323,6 +601,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
o_sf_scale: float = 1.0,
|
o_sf_scale: float = 1.0,
|
||||||
): # Q/K/V arrive already in FP8 via the model-side fused path
|
): # Q/K/V arrive already in FP8 via the model-side fused path
|
||||||
# (prepare_prefill_qkv / pack_prefix_chunk_kv); no quantize here.
|
# (prepare_prefill_qkv / pack_prefix_chunk_kv); no quantize here.
|
||||||
|
# Hybrid MLA models resolve the model-side hook through the outer
|
||||||
|
# HybridLinearAttnBackend, so their fallback MHA path can pass V as a
|
||||||
|
# last-dimension slice of kv_b_proj (stride(-2) > size(-1)). The
|
||||||
|
# TokenSpeed prefill kernel requires dense Q/K/V layouts even though
|
||||||
|
# the public wrapper accepts arbitrary torch tensors.
|
||||||
|
q = q.contiguous()
|
||||||
|
k = k.contiguous()
|
||||||
|
v = v.contiguous()
|
||||||
return tokenspeed_mla.tokenspeed_mla_prefill(
|
return tokenspeed_mla.tokenspeed_mla_prefill(
|
||||||
query=q,
|
query=q,
|
||||||
key=k,
|
key=k,
|
||||||
|
|||||||
@@ -150,6 +150,7 @@ class TRTLLMMLADecodeMetadata:
|
|||||||
cu_seqlens_q: Optional[torch.Tensor] = None
|
cu_seqlens_q: Optional[torch.Tensor] = None
|
||||||
seq_lens_q: Optional[torch.Tensor] = None
|
seq_lens_q: Optional[torch.Tensor] = None
|
||||||
seq_lens_k: Optional[torch.Tensor] = None
|
seq_lens_k: Optional[torch.Tensor] = None
|
||||||
|
global_seq_lens_k: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||||
@@ -382,6 +383,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
|
|
||||||
if forward_mode.is_target_verify():
|
if forward_mode.is_target_verify():
|
||||||
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
|
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
|
||||||
|
metadata.global_seq_lens_k = torch.zeros(
|
||||||
|
(bs,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
elif forward_mode.is_draft_extend_v2():
|
elif forward_mode.is_draft_extend_v2():
|
||||||
num_tokens_per_req = self.num_draft_tokens
|
num_tokens_per_req = self.num_draft_tokens
|
||||||
metadata.max_seq_len_q = num_tokens_per_req
|
metadata.max_seq_len_q = num_tokens_per_req
|
||||||
@@ -423,7 +427,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
|
|
||||||
if forward_mode.is_target_verify():
|
if forward_mode.is_target_verify():
|
||||||
# Intentional int64 -> int32 same-kind out= downcast.
|
# Intentional int64 -> int32 same-kind out= downcast.
|
||||||
torch.add(seq_lens[:bs], self.num_draft_tokens, out=metadata.seq_lens_k)
|
torch.add(
|
||||||
|
seq_lens[:bs],
|
||||||
|
self.num_draft_tokens,
|
||||||
|
out=metadata.global_seq_lens_k,
|
||||||
|
)
|
||||||
|
metadata.seq_lens_k.copy_(metadata.global_seq_lens_k)
|
||||||
seq_lens = metadata.seq_lens_k
|
seq_lens = metadata.seq_lens_k
|
||||||
elif forward_mode.is_draft_extend_v2():
|
elif forward_mode.is_draft_extend_v2():
|
||||||
num_tokens_per_req = self.num_draft_tokens
|
num_tokens_per_req = self.num_draft_tokens
|
||||||
@@ -569,6 +578,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
max_seq = max_seq + self.num_draft_tokens
|
max_seq = max_seq + self.num_draft_tokens
|
||||||
seq_lens = seq_lens + self.num_draft_tokens
|
seq_lens = seq_lens + self.num_draft_tokens
|
||||||
self.forward_decode_metadata.seq_lens_k = seq_lens.to(torch.int32)
|
self.forward_decode_metadata.seq_lens_k = seq_lens.to(torch.int32)
|
||||||
|
self.forward_decode_metadata.global_seq_lens_k = (
|
||||||
|
self.forward_decode_metadata.seq_lens_k
|
||||||
|
)
|
||||||
elif forward_batch.forward_mode.is_draft_extend_v2():
|
elif forward_batch.forward_mode.is_draft_extend_v2():
|
||||||
sum_seq_lens_q = sum(forward_batch.extend_seq_lens_cpu)
|
sum_seq_lens_q = sum(forward_batch.extend_seq_lens_cpu)
|
||||||
max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
|
max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
|
||||||
@@ -951,8 +963,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
q = q.to(self.data_type)
|
q = q.to(self.data_type)
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_target_verify():
|
if forward_batch.forward_mode.is_target_verify():
|
||||||
max_seq_len = (
|
draft_token_num = forward_batch.spec_info.draft_token_num
|
||||||
metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num
|
dcp_enabled = get_parallel().dcp_enabled
|
||||||
|
max_seq_len = metadata.max_seq_len_k + (
|
||||||
|
0 if dcp_enabled else draft_token_num
|
||||||
)
|
)
|
||||||
# For target_verify, all sequences have the same number of draft tokens
|
# For target_verify, all sequences have the same number of draft tokens
|
||||||
q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
|
q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
|
||||||
@@ -1006,6 +1020,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
|
|
||||||
assert kv_cache.dtype == self.data_type
|
assert kv_cache.dtype == self.data_type
|
||||||
|
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
and get_parallel().dcp_enabled
|
||||||
|
):
|
||||||
|
raw_out, lse = self._run_decode_kernel(
|
||||||
|
query=q,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
block_tables=metadata.block_kv_indices,
|
||||||
|
seq_lens=metadata.seq_lens_k,
|
||||||
|
max_seq_len=max_seq_len,
|
||||||
|
layer=layer,
|
||||||
|
causal_seqs=metadata.global_seq_lens_k,
|
||||||
|
cp_world=get_parallel().dcp_size,
|
||||||
|
cp_rank=get_parallel().dcp_rank,
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
output = raw_out.view(
|
||||||
|
bs * draft_token_num,
|
||||||
|
layer.tp_q_head_num,
|
||||||
|
layer.v_head_dim,
|
||||||
|
)
|
||||||
|
lse = lse.view(bs * draft_token_num, layer.tp_q_head_num)
|
||||||
|
dense_q_indptr = torch.arange(
|
||||||
|
0,
|
||||||
|
(bs + 1) * draft_token_num,
|
||||||
|
draft_token_num,
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=q.device,
|
||||||
|
)
|
||||||
|
fixup_zero_kv_rows(
|
||||||
|
output,
|
||||||
|
lse,
|
||||||
|
metadata.seq_lens_k,
|
||||||
|
dense_q_indptr,
|
||||||
|
draft_token_num,
|
||||||
|
)
|
||||||
|
return output.flatten(1), lse
|
||||||
|
|
||||||
raw_out = self._run_decode_kernel(
|
raw_out = self._run_decode_kernel(
|
||||||
query=q,
|
query=q,
|
||||||
kv_cache=kv_cache,
|
kv_cache=kv_cache,
|
||||||
|
|||||||
@@ -3634,6 +3634,12 @@ class HybridLinearKVPool(KVCache):
|
|||||||
def get_kv_size_bytes(self):
|
def get_kv_size_bytes(self):
|
||||||
return self.full_kv_pool.get_kv_size_bytes()
|
return self.full_kv_pool.get_kv_size_bytes()
|
||||||
|
|
||||||
|
def get_kv_buffer_shape(self) -> Tuple[torch.Size, torch.Size]:
|
||||||
|
# Hybrid layer ids are global model-layer ids, while the backing pool
|
||||||
|
# is dense over only full-attention layers. Shape discovery does not
|
||||||
|
# need a global layer lookup, so delegate it to that backing pool.
|
||||||
|
return self.full_kv_pool.get_kv_buffer_shape()
|
||||||
|
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
return self.full_kv_pool.get_contiguous_buf_infos()
|
return self.full_kv_pool.get_contiguous_buf_infos()
|
||||||
|
|
||||||
|
|||||||
@@ -23,8 +23,11 @@ from contextlib import contextmanager
|
|||||||
from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
|
from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
|
||||||
|
|
||||||
from sglang.srt.model_executor.runner.base_runner import BaseRunner
|
from sglang.srt.model_executor.runner.base_runner import BaseRunner
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
from sglang.srt.runtime_context import get_flags
|
||||||
from sglang.srt.utils import require_gathered_buffer
|
from sglang.srt.utils import (
|
||||||
|
get_cuda_graph_batch_size_alignment,
|
||||||
|
get_cuda_graph_max_batch_size,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
@@ -68,22 +71,15 @@ def get_batch_sizes_to_capture(
|
|||||||
capture_bs = list(server_args.cuda_graph_config.decode.bs)
|
capture_bs = list(server_args.cuda_graph_config.decode.bs)
|
||||||
num_max_requests = model_runner.req_to_token_pool.size
|
num_max_requests = model_runner.req_to_token_pool.size
|
||||||
|
|
||||||
mul_base = 1
|
mul_base = get_cuda_graph_batch_size_alignment(server_args)
|
||||||
# TBO splits each request's rows across two micro-batches, so the
|
# TBO splits each request's rows across two micro-batches, so the
|
||||||
# alignment constraint applies per request rather than per token row.
|
# alignment constraint applies per request rather than per token row.
|
||||||
alignment_width = captured_req_width
|
alignment_width = captured_req_width
|
||||||
if server_args.enable_two_batch_overlap:
|
if server_args.enable_two_batch_overlap:
|
||||||
mul_base *= 2
|
|
||||||
alignment_width = 1
|
alignment_width = 1
|
||||||
|
|
||||||
if require_gathered_buffer(server_args):
|
|
||||||
mul_base *= get_parallel().attn_tp_size
|
|
||||||
|
|
||||||
if mul_base % get_parallel().attn_cp_size != 0:
|
|
||||||
mul_base *= get_parallel().attn_cp_size
|
|
||||||
|
|
||||||
# pad `num_max_requests` to avoid being filtered out
|
# pad `num_max_requests` to avoid being filtered out
|
||||||
num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base
|
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests)
|
||||||
if max(capture_bs) > num_max_requests:
|
if max(capture_bs) > num_max_requests:
|
||||||
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
|
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
|
||||||
# is very small. We add more values here to make sure we capture the maximum bs.
|
# is very small. We add more values here to make sure we capture the maximum bs.
|
||||||
|
|||||||
@@ -50,7 +50,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
set_tc_piecewise_forward_context,
|
set_tc_piecewise_forward_context,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_hip
|
||||||
from sglang.srt.utils.common import ceil_align, require_mlp_sync
|
from sglang.srt.utils.common import (
|
||||||
|
ceil_align,
|
||||||
|
get_eager_max_batch_size,
|
||||||
|
require_mlp_sync,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -101,14 +105,12 @@ class EagerRunner(BaseRunner):
|
|||||||
# (expand_for_topk_draft) before the eager fallback.
|
# (expand_for_topk_draft) before the eager fallback.
|
||||||
max_bs *= sa.speculative_eagle_topk
|
max_bs *= sa.speculative_eagle_topk
|
||||||
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
||||||
if require_mlp_sync(sa):
|
max_bs = get_eager_max_batch_size(sa, max_bs)
|
||||||
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
|
||||||
|
|
||||||
max_bs = ceil_align(max_bs, self.attn_tp_size)
|
|
||||||
max_bs = ceil_align(max_bs, get_cp_padding_align_size())
|
|
||||||
prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens())
|
prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens())
|
||||||
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
|
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
|
||||||
if require_mlp_sync(sa):
|
if require_mlp_sync(sa):
|
||||||
|
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||||
|
|
||||||
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
|
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
|
||||||
max_num_token = ceil_align(max_num_token, get_cp_padding_align_size())
|
max_num_token = ceil_align(max_num_token, get_cp_padding_align_size())
|
||||||
self._eager_max_bs = max_bs
|
self._eager_max_bs = max_bs
|
||||||
@@ -261,7 +263,9 @@ class EagerRunner(BaseRunner):
|
|||||||
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
|
|
||||||
if forward_batch.needs_forward_metadata_init():
|
if forward_batch.needs_forward_metadata_init():
|
||||||
if hasattr(model_runner.model, "prepare_context_parallel_metadata_for_dcp"):
|
if model_runner.dcp_size > 1 and hasattr(
|
||||||
|
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
|
||||||
|
):
|
||||||
# prepare kv cache buffer for dcp to gather kv cache
|
# prepare kv cache buffer for dcp to gather kv cache
|
||||||
forward_batch.attn_dcp_metadata = (
|
forward_batch.attn_dcp_metadata = (
|
||||||
model_runner.model.prepare_context_parallel_metadata_for_dcp(
|
model_runner.model.prepare_context_parallel_metadata_for_dcp(
|
||||||
|
|||||||
@@ -88,6 +88,25 @@ class MlaBmmFusionPlan:
|
|||||||
attn_output_buf: torch.Tensor
|
attn_output_buf: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def _is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool:
|
||||||
|
if not get_parallel().dcp_enabled:
|
||||||
|
return False
|
||||||
|
if forward_batch.forward_mode.is_decode():
|
||||||
|
return True
|
||||||
|
if not forward_batch.forward_mode.is_target_verify() or not _is_cuda:
|
||||||
|
return False
|
||||||
|
|
||||||
|
server_args = get_server_args()
|
||||||
|
decode_backend = (
|
||||||
|
server_args.decode_attention_backend or server_args.attention_backend
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
server_args.speculative_algorithm == "DSPARK"
|
||||||
|
and server_args.speculative_attention_mode == "decode"
|
||||||
|
and decode_backend in ("tokenspeed_mla", "cutedsl_mla")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sglang.kernels.ops.gemm import bmm_fp8
|
from sglang.kernels.ops.gemm import bmm_fp8
|
||||||
|
|
||||||
@@ -254,8 +273,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
# weights and skip the per-layer Q all-gather (bf16 decode absorb only).
|
# weights and skip the per-layer Q all-gather (bf16 decode absorb only).
|
||||||
q_replicate_active = (
|
q_replicate_active = (
|
||||||
get_server_args().dcp_replicate_q_proj
|
get_server_args().dcp_replicate_q_proj
|
||||||
and get_parallel().dcp_enabled
|
and _is_dcp_mla_decode_phase(forward_batch)
|
||||||
and forward_batch.forward_mode.is_decode()
|
|
||||||
and not self.use_deep_gemm_bmm
|
and not self.use_deep_gemm_bmm
|
||||||
and self.w_kc_qrep is not None
|
and self.w_kc_qrep is not None
|
||||||
and self.q_b_proj_qrep_weight is not None
|
and self.q_b_proj_qrep_weight is not None
|
||||||
@@ -595,8 +613,8 @@ class DeepseekMLAForwardMixin:
|
|||||||
|
|
||||||
# all_gather q_pe, q_nope_out,take tp8 as an example, q_pe [B, H, ROPE_DIM], q_nope_out [B, H, NOPE_DIM] gathered to [B, H * dcp_world_size, ROPE_DIM] [B, H * dcp_world_size, NOPE_DIM] for decode batch, and all gather k_pe, k_nope for extend batch.
|
# all_gather q_pe, q_nope_out,take tp8 as an example, q_pe [B, H, ROPE_DIM], q_nope_out [B, H, NOPE_DIM] gathered to [B, H * dcp_world_size, ROPE_DIM] [B, H * dcp_world_size, NOPE_DIM] for decode batch, and all gather k_pe, k_nope for extend batch.
|
||||||
if get_parallel().dcp_enabled:
|
if get_parallel().dcp_enabled:
|
||||||
if forward_batch.forward_mode.is_decode() and not q_replicate_active:
|
if _is_dcp_mla_decode_phase(forward_batch):
|
||||||
# if forward_batch.forward_mode is decode, gather q
|
if not q_replicate_active:
|
||||||
q_nope_out, q_pe = all_gather_q_for_mla_decode(
|
q_nope_out, q_pe = all_gather_q_for_mla_decode(
|
||||||
q_nope_out=q_nope_out,
|
q_nope_out=q_nope_out,
|
||||||
q_pe=q_pe,
|
q_pe=q_pe,
|
||||||
@@ -748,10 +766,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
topk_indices=topk_indices,
|
topk_indices=topk_indices,
|
||||||
)
|
)
|
||||||
attn_output = fusion_plan.attn_output_buf
|
attn_output = fusion_plan.attn_output_buf
|
||||||
elif (
|
elif _is_dcp_mla_decode_phase(forward_batch):
|
||||||
forward_batch.forward_mode.is_decode()
|
|
||||||
and get_parallel().dcp_enabled
|
|
||||||
):
|
|
||||||
# set return_lse=True to correct attn_output
|
# set return_lse=True to correct attn_output
|
||||||
attn_output, lse = self.attn_mqa_for_dcp_decode(
|
attn_output, lse = self.attn_mqa_for_dcp_decode(
|
||||||
q_nope_out,
|
q_nope_out,
|
||||||
@@ -825,7 +840,7 @@ class DeepseekMLAForwardMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# correct attn_output with respect to lse from other ranks
|
# correct attn_output with respect to lse from other ranks
|
||||||
if forward_batch.forward_mode.is_decode() and get_parallel().dcp_enabled:
|
if _is_dcp_mla_decode_phase(forward_batch):
|
||||||
attn_output = attn_output.view(
|
attn_output = attn_output.view(
|
||||||
-1,
|
-1,
|
||||||
self.num_local_heads * get_parallel().attn_dcp_size,
|
self.num_local_heads * get_parallel().attn_dcp_size,
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sglang.srt.distributed import (
|
|||||||
tensor_model_parallel_all_reduce,
|
tensor_model_parallel_all_reduce,
|
||||||
)
|
)
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
|
from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
|
||||||
from sglang.srt.layers.layernorm import RMSNorm
|
from sglang.srt.layers.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import (
|
from sglang.srt.layers.linear import (
|
||||||
ColumnParallelBatchedLinear,
|
ColumnParallelBatchedLinear,
|
||||||
@@ -52,6 +53,20 @@ from sglang.srt.utils import make_layers
|
|||||||
from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs
|
from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs
|
||||||
|
|
||||||
|
|
||||||
|
def _get_kda_local_num_heads(num_heads: int, tp_size: int) -> int:
|
||||||
|
if num_heads % tp_size != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"KDA num_heads ({num_heads}) must be divisible by global tp_size ({tp_size})"
|
||||||
|
)
|
||||||
|
return num_heads // tp_size
|
||||||
|
|
||||||
|
|
||||||
|
def _materialize_residual_stream(
|
||||||
|
hidden_states: torch.Tensor, residual: Optional[torch.Tensor]
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return hidden_states if residual is None else hidden_states + residual
|
||||||
|
|
||||||
|
|
||||||
class KimiMoE(nn.Module):
|
class KimiMoE(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -188,8 +203,7 @@ class KimiDeltaAttention(nn.Module):
|
|||||||
self.head_v_dim = config.linear_attn_config["head_dim"]
|
self.head_v_dim = config.linear_attn_config["head_dim"]
|
||||||
self.layer_idx = layer_idx
|
self.layer_idx = layer_idx
|
||||||
self.prefix = prefix
|
self.prefix = prefix
|
||||||
assert self.num_heads % self.tp_size == 0
|
self.local_num_heads = _get_kda_local_num_heads(self.num_heads, self.tp_size)
|
||||||
self.local_num_heads = divide(self.num_heads, self.tp_size)
|
|
||||||
|
|
||||||
projection_size = self.head_dim * self.num_heads
|
projection_size = self.head_dim * self.num_heads
|
||||||
self.conv_size = config.linear_attn_config["short_conv_kernel_size"]
|
self.conv_size = config.linear_attn_config["short_conv_kernel_size"]
|
||||||
@@ -317,9 +331,9 @@ class KimiDeltaAttention(nn.Module):
|
|||||||
|
|
||||||
self.attn = RadixLinearAttention(
|
self.attn = RadixLinearAttention(
|
||||||
layer_id=self.layer_idx,
|
layer_id=self.layer_idx,
|
||||||
num_q_heads=self.num_k_heads // self.attn_tp_size,
|
num_q_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size),
|
||||||
num_k_heads=self.num_k_heads // self.attn_tp_size,
|
num_k_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size),
|
||||||
num_v_heads=self.num_v_heads // self.attn_tp_size,
|
num_v_heads=_get_kda_local_num_heads(self.num_v_heads, self.tp_size),
|
||||||
head_q_dim=self.head_k_dim,
|
head_q_dim=self.head_k_dim,
|
||||||
head_k_dim=self.head_k_dim,
|
head_k_dim=self.head_k_dim,
|
||||||
head_v_dim=self.head_v_dim,
|
head_v_dim=self.head_v_dim,
|
||||||
@@ -519,6 +533,7 @@ class KimiLinearModel(nn.Module):
|
|||||||
self.padding_idx = config.pad_token_id
|
self.padding_idx = config.pad_token_id
|
||||||
self.vocab_size = config.vocab_size
|
self.vocab_size = config.vocab_size
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
|
self.dspark_layers_to_capture: Optional[list[int]] = None
|
||||||
|
|
||||||
if self.pp_group.is_first_rank:
|
if self.pp_group.is_first_rank:
|
||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
@@ -581,7 +596,6 @@ class KimiLinearModel(nn.Module):
|
|||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
# TODO: capture aux hidden states
|
|
||||||
aux_hidden_states = []
|
aux_hidden_states = []
|
||||||
for i in range(self.start_layer, self.end_layer):
|
for i in range(self.start_layer, self.end_layer):
|
||||||
ctx = get_global_expert_distribution_recorder().with_current_layer(i)
|
ctx = get_global_expert_distribution_recorder().with_current_layer(i)
|
||||||
@@ -594,6 +608,13 @@ class KimiLinearModel(nn.Module):
|
|||||||
residual=residual,
|
residual=residual,
|
||||||
zero_allocator=zero_allocator,
|
zero_allocator=zero_allocator,
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
|
self.dspark_layers_to_capture is not None
|
||||||
|
and i in self.dspark_layers_to_capture
|
||||||
|
):
|
||||||
|
aux_hidden_states.append(
|
||||||
|
_materialize_residual_stream(hidden_states, residual)
|
||||||
|
)
|
||||||
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
return PPProxyTensors(
|
return PPProxyTensors(
|
||||||
@@ -609,10 +630,9 @@ class KimiLinearModel(nn.Module):
|
|||||||
else:
|
else:
|
||||||
hidden_states, _ = self.norm(hidden_states, residual)
|
hidden_states, _ = self.norm(hidden_states, residual)
|
||||||
|
|
||||||
if len(aux_hidden_states) == 0:
|
if self.dspark_layers_to_capture is not None:
|
||||||
return hidden_states
|
|
||||||
|
|
||||||
return hidden_states, aux_hidden_states
|
return hidden_states, aux_hidden_states
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
class KimiLinearForCausalLM(nn.Module):
|
class KimiLinearForCausalLM(nn.Module):
|
||||||
@@ -642,6 +662,22 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
self.lm_head = PPMissingLayer()
|
self.lm_head = PPMissingLayer()
|
||||||
logit_scale = getattr(self.config, "logit_scale", 1.0)
|
logit_scale = getattr(self.config, "logit_scale", 1.0)
|
||||||
self.logits_processor = LogitsProcessor(config=config, logit_scale=logit_scale)
|
self.logits_processor = LogitsProcessor(config=config, logit_scale=logit_scale)
|
||||||
|
self.capture_aux_hidden_states = False
|
||||||
|
|
||||||
|
def get_input_embeddings(self):
|
||||||
|
return self.model.embed_tokens
|
||||||
|
|
||||||
|
def set_dspark_layers_to_capture(self, layer_ids: list[int]) -> None:
|
||||||
|
if self.pp_group.world_size > 1:
|
||||||
|
raise NotImplementedError("DSPARK aux hidden capture requires PP=1.")
|
||||||
|
if not self.pp_group.is_last_rank:
|
||||||
|
return
|
||||||
|
if layer_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"DSPARK requires explicit layer_ids for aux hidden capture."
|
||||||
|
)
|
||||||
|
self.capture_aux_hidden_states = True
|
||||||
|
self.model.dspark_layers_to_capture = list(layer_ids)
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
@@ -660,12 +696,47 @@ class KimiLinearForCausalLM(nn.Module):
|
|||||||
pp_proxy_tensors,
|
pp_proxy_tensors,
|
||||||
)
|
)
|
||||||
if self.pp_group.is_last_rank:
|
if self.pp_group.is_last_rank:
|
||||||
|
aux_hidden_states = None
|
||||||
|
if self.capture_aux_hidden_states:
|
||||||
|
hidden_states, aux_hidden_states = hidden_states
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids, hidden_states, self.lm_head, forward_batch
|
input_ids,
|
||||||
|
hidden_states,
|
||||||
|
self.lm_head,
|
||||||
|
forward_batch,
|
||||||
|
aux_hidden_states,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|
||||||
|
def prepare_context_parallel_metadata_for_dcp(
|
||||||
|
self,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
extend_prefix_lens: torch.Tensor,
|
||||||
|
extend_prefix_lens_cpu: torch.Tensor,
|
||||||
|
extend_seq_lens: torch.Tensor,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
seq_lens_sum: int,
|
||||||
|
kv_buffer_shape: torch.Size,
|
||||||
|
kv_cache_dtype,
|
||||||
|
kv_cache_device,
|
||||||
|
create_chunked_prefix_cache_kv_indices_fn,
|
||||||
|
):
|
||||||
|
return prepare_decode_context_parallel_metadata(
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
extend_prefix_lens=extend_prefix_lens,
|
||||||
|
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
|
||||||
|
extend_seq_lens=extend_seq_lens,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
req_to_token=req_to_token,
|
||||||
|
seq_lens_sum=seq_lens_sum,
|
||||||
|
kv_buffer_shape=kv_buffer_shape,
|
||||||
|
kv_cache_dtype=kv_cache_dtype,
|
||||||
|
kv_cache_device=kv_cache_device,
|
||||||
|
create_chunked_prefix_cache_kv_indices_fn=create_chunked_prefix_cache_kv_indices_fn,
|
||||||
|
)
|
||||||
|
|
||||||
def _is_non_local_pp_weight(self, name: str) -> bool:
|
def _is_non_local_pp_weight(self, name: str) -> bool:
|
||||||
if self.pp_group.world_size == 1:
|
if self.pp_group.world_size == 1:
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -1038,14 +1038,6 @@ class ServerArgs:
|
|||||||
),
|
),
|
||||||
NS("parallel"),
|
NS("parallel"),
|
||||||
] = 1
|
] = 1
|
||||||
dcp_size: A[
|
|
||||||
int,
|
|
||||||
Arg(
|
|
||||||
help="The decode context parallelism size.",
|
|
||||||
aliases=["--decode-context-parallel-size"],
|
|
||||||
),
|
|
||||||
NS("parallel"),
|
|
||||||
] = 1
|
|
||||||
dwdp_size: A[
|
dwdp_size: A[
|
||||||
int,
|
int,
|
||||||
Arg(
|
Arg(
|
||||||
@@ -1064,20 +1056,24 @@ class ServerArgs:
|
|||||||
"combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires "
|
"combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires "
|
||||||
"SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).",
|
"SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).",
|
||||||
choices=["ag_rs", "a2a", "fi_a2a"],
|
choices=["ag_rs", "a2a", "fi_a2a"],
|
||||||
|
resolvable=True,
|
||||||
),
|
),
|
||||||
NS("parallel"),
|
NS("parallel"),
|
||||||
] = "ag_rs"
|
] = "ag_rs"
|
||||||
dcp_replicate_q_proj: A[
|
dcp_replicate_q_proj: A[
|
||||||
bool,
|
Optional[bool],
|
||||||
Arg(
|
Arg(
|
||||||
help="For MLA decode context parallelism with the a2a/fi_a2a "
|
help="For MLA decode context parallelism with the a2a/fi_a2a "
|
||||||
"backend: replicate the Q projection so each DCP rank computes the "
|
"backend: replicate the Q projection so each DCP rank computes the "
|
||||||
"full-head query locally (redundant projection compute), eliminating "
|
"full-head query locally (redundant projection compute), eliminating "
|
||||||
"the per-layer head-dim all-gather of Q. Trades a small amount of "
|
"the per-layer head-dim all-gather of Q. Trades a small amount of "
|
||||||
"extra GEMM for one fewer collective per layer.",
|
"extra GEMM for one fewer collective per layer. Use "
|
||||||
|
"--no-dcp-replicate-q-proj to disable the model-specific default.",
|
||||||
|
action=argparse.BooleanOptionalAction,
|
||||||
|
resolvable=True,
|
||||||
),
|
),
|
||||||
NS("parallel"),
|
NS("parallel"),
|
||||||
] = False
|
] = None
|
||||||
enable_prefill_cp: A[
|
enable_prefill_cp: A[
|
||||||
bool,
|
bool,
|
||||||
"Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.",
|
"Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.",
|
||||||
@@ -3763,12 +3759,34 @@ class ServerArgs:
|
|||||||
return
|
return
|
||||||
elif is_cuda():
|
elif is_cuda():
|
||||||
if self.speculative_algorithm is not None:
|
if self.speculative_algorithm is not None:
|
||||||
|
model_arches = self.get_model_config().hf_config.architectures
|
||||||
|
decode_backend = self.decode_attention_backend or self.attention_backend
|
||||||
|
kimi_linear_dspark = (
|
||||||
|
self.speculative_algorithm == "DSPARK"
|
||||||
|
and "KimiLinearForCausalLM" in model_arches
|
||||||
|
and self.speculative_attention_mode == "decode"
|
||||||
|
and decode_backend in ("tokenspeed_mla", "cutedsl_mla")
|
||||||
|
)
|
||||||
|
if kimi_linear_dspark:
|
||||||
|
ragged_verify_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get()
|
||||||
|
if ragged_verify_mode != "static":
|
||||||
|
raise ValueError(
|
||||||
|
"Kimi Linear DCP + DSPARK currently requires "
|
||||||
|
"SGLANG_RAGGED_VERIFY_MODE=static, but got "
|
||||||
|
f"{ragged_verify_mode!r}."
|
||||||
|
)
|
||||||
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Decode context parallel (--dcp-size / "
|
"Decode context parallel (--dcp-size / "
|
||||||
"--decode-context-parallel-size > 1) on CUDA platform "
|
"--decode-context-parallel-size > 1) with speculative "
|
||||||
"does not support any speculative algorithm, but got "
|
"decoding on CUDA is supported only for Kimi Linear + "
|
||||||
f"dcp_size={self.dcp_size} on a CUDA platform with "
|
"DSPARK + --speculative-attention-mode decode + "
|
||||||
"speculative decoding enabled."
|
"tokenspeed_mla, or experimental cutedsl_mla, but got "
|
||||||
|
f"architectures={model_arches}, "
|
||||||
|
f"speculative_algorithm={self.speculative_algorithm!r}, "
|
||||||
|
"speculative_attention_mode="
|
||||||
|
f"{self.speculative_attention_mode!r}, "
|
||||||
|
f"decode_attention_backend={decode_backend!r}."
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -239,7 +239,18 @@ class DraftBackendFactory:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _create_cutedsl_mla_decode_backend(self):
|
def _create_cutedsl_mla_decode_backend(self):
|
||||||
return self._create_trtllm_mla_decode_backend(backend="cute-dsl")
|
if not self.draft_model_runner.use_mla_backend:
|
||||||
|
raise ValueError(
|
||||||
|
"cutedsl_mla backend requires MLA model (use_mla_backend=True)."
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention.cutedsl_mla_backend import (
|
||||||
|
CuteDslMLAMultiStepDraftBackend,
|
||||||
|
)
|
||||||
|
|
||||||
|
return CuteDslMLAMultiStepDraftBackend(
|
||||||
|
self.draft_model_runner, self.topk, self.speculative_num_steps
|
||||||
|
)
|
||||||
|
|
||||||
def _create_tokenspeed_mla_decode_backend(self):
|
def _create_tokenspeed_mla_decode_backend(self):
|
||||||
if not self.draft_model_runner.use_mla_backend:
|
if not self.draft_model_runner.use_mla_backend:
|
||||||
|
|||||||
@@ -3608,6 +3608,31 @@ def require_mlp_sync(server_args: ServerArgs):
|
|||||||
return server_args.enable_dp_attention or require_gathered_buffer(server_args)
|
return server_args.enable_dp_attention or require_gathered_buffer(server_args)
|
||||||
|
|
||||||
|
|
||||||
|
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
|
||||||
|
alignment = 1
|
||||||
|
if server_args.enable_two_batch_overlap:
|
||||||
|
alignment *= 2
|
||||||
|
if require_gathered_buffer(server_args):
|
||||||
|
alignment *= get_parallel().attn_tp_size
|
||||||
|
if alignment % get_parallel().attn_cp_size != 0:
|
||||||
|
alignment *= get_parallel().attn_cp_size
|
||||||
|
return alignment
|
||||||
|
|
||||||
|
|
||||||
|
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
|
||||||
|
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args))
|
||||||
|
|
||||||
|
|
||||||
|
def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
|
||||||
|
if not require_mlp_sync(server_args):
|
||||||
|
return max_batch_size
|
||||||
|
|
||||||
|
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
|
||||||
|
|
||||||
|
max_batch_size = ceil_align(max_batch_size, get_parallel().attn_tp_size)
|
||||||
|
return ceil_align(max_batch_size, get_cp_padding_align_size())
|
||||||
|
|
||||||
|
|
||||||
def find_local_repo_dir(repo_id: str, revision: Optional[str] = None) -> Optional[str]:
|
def find_local_repo_dir(repo_id: str, revision: Optional[str] = None) -> Optional[str]:
|
||||||
import huggingface_hub as hf
|
import huggingface_hub as hf
|
||||||
|
|
||||||
|
|||||||
@@ -12,12 +12,19 @@ Usage:
|
|||||||
python test_dcp_layout_unit.py
|
python test_dcp_layout_unit.py
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import math
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
|
||||||
|
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
|
||||||
|
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
|
|
||||||
@@ -36,7 +43,7 @@ def _legacy_inplace_formula(length: int, n: int, rank: int) -> int:
|
|||||||
return (length - rank - 1) // n + 1
|
return (length - rank - 1) // n + 1
|
||||||
|
|
||||||
|
|
||||||
class TestGetDcpLens(unittest.TestCase):
|
class TestGetDcpLens(CustomTestCase):
|
||||||
def test_start_none_matches_owner_count(self):
|
def test_start_none_matches_owner_count(self):
|
||||||
for n in DCP_SIZES:
|
for n in DCP_SIZES:
|
||||||
for rank in range(n):
|
for rank in range(n):
|
||||||
@@ -86,6 +93,156 @@ class TestGetDcpLens(unittest.TestCase):
|
|||||||
lens = torch.tensor(LENS, dtype=torch.int32)
|
lens = torch.tensor(LENS, dtype=torch.int32)
|
||||||
self.assertTrue(torch.equal(get_dcp_lens(lens, 1, 0), lens))
|
self.assertTrue(torch.equal(get_dcp_lens(lens, 1, 0), lens))
|
||||||
|
|
||||||
|
def test_paged_allocator_exposes_dcp_virtual_capacity(self):
|
||||||
|
real_kv_size = 1024
|
||||||
|
dcp_size = 4
|
||||||
|
physical_page_size = 64
|
||||||
|
allocator = PagedTokenToKVPoolAllocator(
|
||||||
|
size=real_kv_size * dcp_size,
|
||||||
|
page_size=physical_page_size * dcp_size,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device="cpu",
|
||||||
|
kvcache=object(),
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
allocations = [allocator.alloc(physical_page_size * dcp_size) for _ in range(4)]
|
||||||
|
self.assertTrue(all(indices is not None for indices in allocations))
|
||||||
|
virtual_indices = torch.cat(allocations)
|
||||||
|
|
||||||
|
self.assertEqual(allocator.size, real_kv_size * dcp_size)
|
||||||
|
self.assertEqual(allocator.page_size, physical_page_size * dcp_size)
|
||||||
|
self.assertEqual(allocator.num_pages, real_kv_size // physical_page_size)
|
||||||
|
self.assertEqual(
|
||||||
|
len(torch.unique(virtual_indices // dcp_size)),
|
||||||
|
len(virtual_indices) // dcp_size,
|
||||||
|
)
|
||||||
|
self.assertLess(
|
||||||
|
int((virtual_indices // dcp_size).max()),
|
||||||
|
real_kv_size + physical_page_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_configurator_scales_only_the_virtual_dcp_allocator(self):
|
||||||
|
physical_kv_size = 1024
|
||||||
|
physical_page_size = 64
|
||||||
|
physical_kv_cache = SimpleNamespace(
|
||||||
|
size=physical_kv_size,
|
||||||
|
page_size=physical_page_size,
|
||||||
|
)
|
||||||
|
sizes = SimpleNamespace(
|
||||||
|
max_total_num_tokens=physical_kv_size,
|
||||||
|
full_max_total_num_tokens=None,
|
||||||
|
swa_max_total_num_tokens=None,
|
||||||
|
)
|
||||||
|
allocators = {}
|
||||||
|
|
||||||
|
for dcp_size in (1, 4):
|
||||||
|
configurator = SimpleNamespace(
|
||||||
|
server_args=SimpleNamespace(
|
||||||
|
disaggregation_mode="null",
|
||||||
|
enable_hisparse=False,
|
||||||
|
page_size=physical_page_size,
|
||||||
|
dcp_size=dcp_size,
|
||||||
|
),
|
||||||
|
hybrid_gdn_config=None,
|
||||||
|
is_hybrid_swa=False,
|
||||||
|
kv_cache_dtype=torch.bfloat16,
|
||||||
|
device="cpu",
|
||||||
|
is_draft_worker=False,
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.mem_cache.kv_cache_configurator.current_platform.is_out_of_tree",
|
||||||
|
return_value=False,
|
||||||
|
):
|
||||||
|
allocators[dcp_size] = (
|
||||||
|
KVCacheConfigurator._build_token_to_kv_pool_allocator(
|
||||||
|
configurator,
|
||||||
|
sizes=sizes,
|
||||||
|
token_to_kv_pool=physical_kv_cache,
|
||||||
|
is_dsv4_model=False,
|
||||||
|
req_to_token_pool=object(),
|
||||||
|
token_to_kv_pool_allocator=None,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
dcp1_allocator = allocators[1]
|
||||||
|
dcp4_allocator = allocators[4]
|
||||||
|
self.assertIs(dcp1_allocator.get_kvcache(), physical_kv_cache)
|
||||||
|
self.assertIs(dcp4_allocator.get_kvcache(), physical_kv_cache)
|
||||||
|
self.assertEqual(dcp1_allocator.size, 1024)
|
||||||
|
self.assertEqual(dcp1_allocator.page_size, 64)
|
||||||
|
self.assertEqual(dcp1_allocator.num_pages, 16)
|
||||||
|
self.assertEqual(dcp4_allocator.size, 4096)
|
||||||
|
self.assertEqual(dcp4_allocator.page_size, 256)
|
||||||
|
self.assertEqual(dcp4_allocator.num_pages, 16)
|
||||||
|
|
||||||
|
def test_live_cell_and_page_ownership_formulas(self):
|
||||||
|
dcp_size = 4
|
||||||
|
physical_page_size = 64
|
||||||
|
ragged_lengths = (0, 1, 2, 3, 4, 63, 64, 65, 255, 256, 257, 515)
|
||||||
|
|
||||||
|
per_rank_counts = []
|
||||||
|
for rank in range(dcp_size):
|
||||||
|
expected_counts = [
|
||||||
|
length // dcp_size + int(rank < length % dcp_size)
|
||||||
|
for length in ragged_lengths
|
||||||
|
]
|
||||||
|
actual_counts = [
|
||||||
|
_owner_count(length, dcp_size, rank, 0) for length in ragged_lengths
|
||||||
|
]
|
||||||
|
self.assertEqual(actual_counts, expected_counts)
|
||||||
|
per_rank_counts.append(sum(actual_counts))
|
||||||
|
|
||||||
|
allocated_pages = [
|
||||||
|
math.ceil(length / (physical_page_size * dcp_size))
|
||||||
|
for length in ragged_lengths
|
||||||
|
]
|
||||||
|
active_pages = [
|
||||||
|
math.ceil(count / physical_page_size) for count in actual_counts
|
||||||
|
]
|
||||||
|
self.assertTrue(
|
||||||
|
all(
|
||||||
|
active <= allocated
|
||||||
|
for active, allocated in zip(active_pages, allocated_pages)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
all(
|
||||||
|
allocated - active <= 1
|
||||||
|
for active, allocated in zip(active_pages, allocated_pages)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(sum(per_rank_counts), sum(ragged_lengths))
|
||||||
|
|
||||||
|
aligned_lengths = (256, 512, 768, 1024)
|
||||||
|
full_replica_cells = sum(aligned_lengths)
|
||||||
|
full_replica_pages = sum(
|
||||||
|
length // physical_page_size for length in aligned_lengths
|
||||||
|
)
|
||||||
|
for rank in range(dcp_size):
|
||||||
|
local_cells = sum(
|
||||||
|
_owner_count(length, dcp_size, rank, 0) for length in aligned_lengths
|
||||||
|
)
|
||||||
|
local_pages = sum(
|
||||||
|
math.ceil(_owner_count(length, dcp_size, rank, 0) / physical_page_size)
|
||||||
|
for length in aligned_lengths
|
||||||
|
)
|
||||||
|
self.assertEqual(local_cells * dcp_size, full_replica_cells)
|
||||||
|
self.assertEqual(local_pages * dcp_size, full_replica_pages)
|
||||||
|
|
||||||
|
def test_hybrid_pool_reports_the_backing_attention_shape(self):
|
||||||
|
pool = object.__new__(HybridLinearKVPool)
|
||||||
|
pool.start_layer = 0
|
||||||
|
pool.layer_transfer_counter = None
|
||||||
|
pool.full_attention_layer_id_mapping = {3: 0, 7: 1}
|
||||||
|
pool.full_kv_pool = MagicMock()
|
||||||
|
expected = (torch.Size([1024, 1, 576]), torch.Size([1024, 1, 576]))
|
||||||
|
pool.full_kv_pool.get_kv_buffer_shape.return_value = expected
|
||||||
|
|
||||||
|
self.assertEqual(pool.get_kv_buffer_shape(), expected)
|
||||||
|
pool.full_kv_pool.get_kv_buffer_shape.assert_called_once_with()
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,129 @@
|
|||||||
|
"""Four-Blackwell acceptance coverage for Kimi Linear TokenSpeed MLA DCP.
|
||||||
|
|
||||||
|
The captured-shape and eager-shape requests deliberately straddle
|
||||||
|
``--cuda-graph-max-bs-decode=64``. This guards both the regular CUDA graph
|
||||||
|
decode path and the full-capacity eager DCP LSE scratch-buffer path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=900, stage="base-c", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct"
|
||||||
|
|
||||||
|
|
||||||
|
def _has_four_blackwell_gpus() -> bool:
|
||||||
|
if not torch.cuda.is_available() or torch.cuda.device_count() < 4:
|
||||||
|
return False
|
||||||
|
return all(
|
||||||
|
torch.cuda.get_device_capability(device_index) >= (10, 0)
|
||||||
|
for device_index in range(4)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(
|
||||||
|
_has_four_blackwell_gpus(),
|
||||||
|
"TokenSpeed MLA DCP acceptance requires four Blackwell GPUs",
|
||||||
|
)
|
||||||
|
class TestKimiLinearDCP4(GSM8KMixin, CustomTestCase):
|
||||||
|
model = KIMI_LINEAR_MODEL
|
||||||
|
base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
gsm8k_score_threshold = 0.90
|
||||||
|
gsm8k_num_examples = 200
|
||||||
|
# Keep accuracy evaluation within the captured decode batch sizes so its
|
||||||
|
# score is batch-invariant. The separate smoke test still exercises the
|
||||||
|
# eager path with batch size 65.
|
||||||
|
gsm8k_num_threads = 4
|
||||||
|
gsm8k_num_shots = 5
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.process = None
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5,
|
||||||
|
other_args=[
|
||||||
|
"--tp-size",
|
||||||
|
"4",
|
||||||
|
"--dcp-size",
|
||||||
|
"4",
|
||||||
|
"--attention-backend",
|
||||||
|
"tokenspeed_mla",
|
||||||
|
"--kv-cache-dtype",
|
||||||
|
"fp8_e4m3",
|
||||||
|
"--dcp-comm-backend",
|
||||||
|
"a2a",
|
||||||
|
"--dcp-replicate-q-proj",
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--random-seed",
|
||||||
|
"0",
|
||||||
|
"--dtype",
|
||||||
|
"bfloat16",
|
||||||
|
"--cuda-graph-max-bs-decode",
|
||||||
|
"64",
|
||||||
|
"--cuda-graph-backend-prefill",
|
||||||
|
"disabled",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.80",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process:
|
||||||
|
kill_process_tree(cls.process.pid, wait_timeout=60)
|
||||||
|
|
||||||
|
def _assert_batch_completes(self, batch_size: int):
|
||||||
|
prompts = [
|
||||||
|
f"Reply with one short word for request {index}: the sky is"
|
||||||
|
for index in range(batch_size)
|
||||||
|
]
|
||||||
|
response = requests.post(
|
||||||
|
self.base_url + "/generate",
|
||||||
|
json={
|
||||||
|
"text": prompts,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0,
|
||||||
|
"max_new_tokens": 8,
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
timeout=180,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
outputs = response.json()
|
||||||
|
self.assertIsInstance(outputs, list)
|
||||||
|
self.assertEqual(len(outputs), batch_size)
|
||||||
|
for output in outputs:
|
||||||
|
self.assertTrue(output["text"].strip())
|
||||||
|
self.assertGreater(output["meta_info"]["completion_tokens"], 0)
|
||||||
|
|
||||||
|
def test_decode_cuda_graph_and_eager_batch(self):
|
||||||
|
# Batch two replays a captured shape; batch 65 is above the configured
|
||||||
|
# regular CUDA graph maximum and therefore exercises eager decode.
|
||||||
|
self._assert_batch_completes(2)
|
||||||
|
self._assert_batch_completes(2)
|
||||||
|
self._assert_batch_completes(65)
|
||||||
|
|
||||||
|
def test_physical_capacity_sanity(self):
|
||||||
|
response = requests.get(self.base_url + "/server_info", timeout=30)
|
||||||
|
response.raise_for_status()
|
||||||
|
self.assertGreater(response.json()["max_total_num_tokens"], 0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -71,6 +71,8 @@ class TestModelOverridableWhitelist(CustomTestCase):
|
|||||||
"ep_size",
|
"ep_size",
|
||||||
"moe_dense_tp_size",
|
"moe_dense_tp_size",
|
||||||
"attn_cp_size",
|
"attn_cp_size",
|
||||||
|
"dcp_comm_backend",
|
||||||
|
"dcp_replicate_q_proj",
|
||||||
"disable_overlap_schedule",
|
"disable_overlap_schedule",
|
||||||
"uses_mamba_radix_cache",
|
"uses_mamba_radix_cache",
|
||||||
"mamba_radix_cache_strategy",
|
"mamba_radix_cache_strategy",
|
||||||
|
|||||||
Reference in New Issue
Block a user