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
|
||||
# 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
|
||||
def create_dcp_kv_indices(
|
||||
kv_indptr,
|
||||
|
||||
@@ -89,9 +89,9 @@ def create_tokenspeed_mla_backend(runner):
|
||||
def create_cutedsl_mla_backend(runner):
|
||||
if not runner.use_mla_backend:
|
||||
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")
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
get_cuda_graph_max_batch_size,
|
||||
get_int_env_var,
|
||||
is_flashinfer_available,
|
||||
is_sm100_supported,
|
||||
next_power_of_2,
|
||||
require_gathered_buffer,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -63,18 +63,6 @@ if TYPE_CHECKING:
|
||||
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():
|
||||
torch._logging.set_logs(dynamo=logging.ERROR)
|
||||
torch._dynamo.config.suppress_errors = True
|
||||
@@ -447,7 +435,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
)
|
||||
else:
|
||||
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
|
||||
)
|
||||
if kv_indptr_buf is None:
|
||||
@@ -2230,7 +2218,7 @@ class FlashInferMultiStepDraftBackend:
|
||||
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
|
||||
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
|
||||
)
|
||||
self.kv_indptr = torch.zeros(
|
||||
|
||||
@@ -22,10 +22,10 @@ from __future__ import annotations
|
||||
|
||||
"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell.
|
||||
|
||||
Subclasses :class:`TRTLLMMLABackend` and overrides only ``_run_decode_kernel``
|
||||
and ``_run_prefill_kernel``. All metadata, KV-cache layout, CUDA-graph
|
||||
plumbing, FP8 quantize/rope, draft-extend padding, and chunked-prefix
|
||||
dispatch are inherited unchanged from the parent.
|
||||
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation and
|
||||
prefill plumbing. Decode-context parallelism is implemented here because the
|
||||
TokenSpeed decode kernel natively accepts CP rank/world metadata and returns
|
||||
the partial log-sum-exp needed by the cross-rank merge.
|
||||
"""
|
||||
|
||||
import logging
|
||||
@@ -34,14 +34,28 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
|
||||
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 (
|
||||
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.srt.layers.attention.trtllm_mla_backend import (
|
||||
TRTLLMMLABackend,
|
||||
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
|
||||
|
||||
if is_flashinfer_available():
|
||||
@@ -273,6 +287,162 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
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(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
@@ -281,7 +451,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
seq_lens: torch.Tensor,
|
||||
max_seq_len: int,
|
||||
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)
|
||||
if k_scale is None:
|
||||
k_scale = 1.0
|
||||
@@ -303,8 +478,111 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
softmax_scale=softmax_scale,
|
||||
output_scale=output_scale,
|
||||
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(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -323,6 +601,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
||||
o_sf_scale: float = 1.0,
|
||||
): # Q/K/V arrive already in FP8 via the model-side fused path
|
||||
# (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(
|
||||
query=q,
|
||||
key=k,
|
||||
|
||||
@@ -150,6 +150,7 @@ class TRTLLMMLADecodeMetadata:
|
||||
cu_seqlens_q: Optional[torch.Tensor] = None
|
||||
seq_lens_q: Optional[torch.Tensor] = None
|
||||
seq_lens_k: Optional[torch.Tensor] = None
|
||||
global_seq_lens_k: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
@@ -382,6 +383,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
|
||||
if forward_mode.is_target_verify():
|
||||
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():
|
||||
num_tokens_per_req = self.num_draft_tokens
|
||||
metadata.max_seq_len_q = num_tokens_per_req
|
||||
@@ -423,7 +427,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
|
||||
if forward_mode.is_target_verify():
|
||||
# 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
|
||||
elif forward_mode.is_draft_extend_v2():
|
||||
num_tokens_per_req = self.num_draft_tokens
|
||||
@@ -569,6 +578,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
max_seq = max_seq + 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.global_seq_lens_k = (
|
||||
self.forward_decode_metadata.seq_lens_k
|
||||
)
|
||||
elif forward_batch.forward_mode.is_draft_extend_v2():
|
||||
sum_seq_lens_q = sum(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)
|
||||
|
||||
if forward_batch.forward_mode.is_target_verify():
|
||||
max_seq_len = (
|
||||
metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num
|
||||
draft_token_num = 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
|
||||
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
|
||||
|
||||
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(
|
||||
query=q,
|
||||
kv_cache=kv_cache,
|
||||
|
||||
@@ -3634,6 +3634,12 @@ class HybridLinearKVPool(KVCache):
|
||||
def get_kv_size_bytes(self):
|
||||
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):
|
||||
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 sglang.srt.model_executor.runner.base_runner import BaseRunner
|
||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
||||
from sglang.srt.utils import require_gathered_buffer
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.utils import (
|
||||
get_cuda_graph_batch_size_alignment,
|
||||
get_cuda_graph_max_batch_size,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
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)
|
||||
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
|
||||
# alignment constraint applies per request rather than per token row.
|
||||
alignment_width = captured_req_width
|
||||
if server_args.enable_two_batch_overlap:
|
||||
mul_base *= 2
|
||||
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
|
||||
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:
|
||||
# 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.
|
||||
|
||||
@@ -50,7 +50,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
set_tc_piecewise_forward_context,
|
||||
)
|
||||
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__)
|
||||
|
||||
@@ -101,14 +105,12 @@ class EagerRunner(BaseRunner):
|
||||
# (expand_for_topk_draft) before the eager fallback.
|
||||
max_bs *= sa.speculative_eagle_topk
|
||||
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
|
||||
if require_mlp_sync(sa):
|
||||
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())
|
||||
max_bs = get_eager_max_batch_size(sa, max_bs)
|
||||
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)
|
||||
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, get_cp_padding_align_size())
|
||||
self._eager_max_bs = max_bs
|
||||
@@ -261,7 +263,9 @@ class EagerRunner(BaseRunner):
|
||||
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||
|
||||
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
|
||||
forward_batch.attn_dcp_metadata = (
|
||||
model_runner.model.prepare_context_parallel_metadata_for_dcp(
|
||||
|
||||
@@ -88,6 +88,25 @@ class MlaBmmFusionPlan:
|
||||
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:
|
||||
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).
|
||||
q_replicate_active = (
|
||||
get_server_args().dcp_replicate_q_proj
|
||||
and get_parallel().dcp_enabled
|
||||
and forward_batch.forward_mode.is_decode()
|
||||
and _is_dcp_mla_decode_phase(forward_batch)
|
||||
and not self.use_deep_gemm_bmm
|
||||
and self.w_kc_qrep is not None
|
||||
and self.q_b_proj_qrep_weight is not None
|
||||
@@ -595,12 +613,12 @@ 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.
|
||||
if get_parallel().dcp_enabled:
|
||||
if forward_batch.forward_mode.is_decode() and not q_replicate_active:
|
||||
# if forward_batch.forward_mode is decode, gather q
|
||||
q_nope_out, q_pe = all_gather_q_for_mla_decode(
|
||||
q_nope_out=q_nope_out,
|
||||
q_pe=q_pe,
|
||||
)
|
||||
if _is_dcp_mla_decode_phase(forward_batch):
|
||||
if not q_replicate_active:
|
||||
q_nope_out, q_pe = all_gather_q_for_mla_decode(
|
||||
q_nope_out=q_nope_out,
|
||||
q_pe=q_pe,
|
||||
)
|
||||
elif forward_batch.forward_mode.is_extend():
|
||||
# for extend, gather kv
|
||||
all_gather_kv_cache_for_mla_extend(
|
||||
@@ -748,10 +766,7 @@ class DeepseekMLAForwardMixin:
|
||||
topk_indices=topk_indices,
|
||||
)
|
||||
attn_output = fusion_plan.attn_output_buf
|
||||
elif (
|
||||
forward_batch.forward_mode.is_decode()
|
||||
and get_parallel().dcp_enabled
|
||||
):
|
||||
elif _is_dcp_mla_decode_phase(forward_batch):
|
||||
# set return_lse=True to correct attn_output
|
||||
attn_output, lse = self.attn_mqa_for_dcp_decode(
|
||||
q_nope_out,
|
||||
@@ -825,7 +840,7 @@ class DeepseekMLAForwardMixin:
|
||||
)
|
||||
|
||||
# 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(
|
||||
-1,
|
||||
self.num_local_heads * get_parallel().attn_dcp_size,
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.distributed import (
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
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.linear import (
|
||||
ColumnParallelBatchedLinear,
|
||||
@@ -52,6 +53,20 @@ from sglang.srt.utils import make_layers
|
||||
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):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -188,8 +203,7 @@ class KimiDeltaAttention(nn.Module):
|
||||
self.head_v_dim = config.linear_attn_config["head_dim"]
|
||||
self.layer_idx = layer_idx
|
||||
self.prefix = prefix
|
||||
assert self.num_heads % self.tp_size == 0
|
||||
self.local_num_heads = divide(self.num_heads, self.tp_size)
|
||||
self.local_num_heads = _get_kda_local_num_heads(self.num_heads, self.tp_size)
|
||||
|
||||
projection_size = self.head_dim * self.num_heads
|
||||
self.conv_size = config.linear_attn_config["short_conv_kernel_size"]
|
||||
@@ -317,9 +331,9 @@ class KimiDeltaAttention(nn.Module):
|
||||
|
||||
self.attn = RadixLinearAttention(
|
||||
layer_id=self.layer_idx,
|
||||
num_q_heads=self.num_k_heads // self.attn_tp_size,
|
||||
num_k_heads=self.num_k_heads // self.attn_tp_size,
|
||||
num_v_heads=self.num_v_heads // self.attn_tp_size,
|
||||
num_q_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size),
|
||||
num_k_heads=_get_kda_local_num_heads(self.num_k_heads, self.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_k_dim=self.head_k_dim,
|
||||
head_v_dim=self.head_v_dim,
|
||||
@@ -519,6 +533,7 @@ class KimiLinearModel(nn.Module):
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.vocab_size = config.vocab_size
|
||||
self.pp_group = get_pp_group()
|
||||
self.dspark_layers_to_capture: Optional[list[int]] = None
|
||||
|
||||
if self.pp_group.is_first_rank:
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
@@ -581,7 +596,6 @@ class KimiLinearModel(nn.Module):
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
# TODO: capture aux hidden states
|
||||
aux_hidden_states = []
|
||||
for i in range(self.start_layer, self.end_layer):
|
||||
ctx = get_global_expert_distribution_recorder().with_current_layer(i)
|
||||
@@ -594,6 +608,13 @@ class KimiLinearModel(nn.Module):
|
||||
residual=residual,
|
||||
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:
|
||||
return PPProxyTensors(
|
||||
@@ -609,10 +630,9 @@ class KimiLinearModel(nn.Module):
|
||||
else:
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
|
||||
if len(aux_hidden_states) == 0:
|
||||
return hidden_states
|
||||
|
||||
return hidden_states, aux_hidden_states
|
||||
if self.dspark_layers_to_capture is not None:
|
||||
return hidden_states, aux_hidden_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
class KimiLinearForCausalLM(nn.Module):
|
||||
@@ -642,6 +662,22 @@ class KimiLinearForCausalLM(nn.Module):
|
||||
self.lm_head = PPMissingLayer()
|
||||
logit_scale = getattr(self.config, "logit_scale", 1.0)
|
||||
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()
|
||||
def forward(
|
||||
@@ -660,12 +696,47 @@ class KimiLinearForCausalLM(nn.Module):
|
||||
pp_proxy_tensors,
|
||||
)
|
||||
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(
|
||||
input_ids, hidden_states, self.lm_head, forward_batch
|
||||
input_ids,
|
||||
hidden_states,
|
||||
self.lm_head,
|
||||
forward_batch,
|
||||
aux_hidden_states,
|
||||
)
|
||||
else:
|
||||
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:
|
||||
if self.pp_group.world_size == 1:
|
||||
return False
|
||||
|
||||
@@ -1038,14 +1038,6 @@ class ServerArgs:
|
||||
),
|
||||
NS("parallel"),
|
||||
] = 1
|
||||
dcp_size: A[
|
||||
int,
|
||||
Arg(
|
||||
help="The decode context parallelism size.",
|
||||
aliases=["--decode-context-parallel-size"],
|
||||
),
|
||||
NS("parallel"),
|
||||
] = 1
|
||||
dwdp_size: A[
|
||||
int,
|
||||
Arg(
|
||||
@@ -1064,20 +1056,24 @@ class ServerArgs:
|
||||
"combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires "
|
||||
"SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).",
|
||||
choices=["ag_rs", "a2a", "fi_a2a"],
|
||||
resolvable=True,
|
||||
),
|
||||
NS("parallel"),
|
||||
] = "ag_rs"
|
||||
dcp_replicate_q_proj: A[
|
||||
bool,
|
||||
Optional[bool],
|
||||
Arg(
|
||||
help="For MLA decode context parallelism with the a2a/fi_a2a "
|
||||
"backend: replicate the Q projection so each DCP rank computes the "
|
||||
"full-head query locally (redundant projection compute), eliminating "
|
||||
"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"),
|
||||
] = False
|
||||
] = None
|
||||
enable_prefill_cp: A[
|
||||
bool,
|
||||
"Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.",
|
||||
@@ -3763,13 +3759,35 @@ class ServerArgs:
|
||||
return
|
||||
elif is_cuda():
|
||||
if self.speculative_algorithm is not None:
|
||||
raise ValueError(
|
||||
"Decode context parallel (--dcp-size / "
|
||||
"--decode-context-parallel-size > 1) on CUDA platform "
|
||||
"does not support any speculative algorithm, but got "
|
||||
f"dcp_size={self.dcp_size} on a CUDA platform with "
|
||||
"speculative decoding enabled."
|
||||
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(
|
||||
"Decode context parallel (--dcp-size / "
|
||||
"--decode-context-parallel-size > 1) with speculative "
|
||||
"decoding on CUDA is supported only for Kimi Linear + "
|
||||
"DSPARK + --speculative-attention-mode decode + "
|
||||
"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:
|
||||
raise ValueError(
|
||||
"Decode context parallel (--dcp-size / "
|
||||
|
||||
@@ -239,7 +239,18 @@ class DraftBackendFactory:
|
||||
)
|
||||
|
||||
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):
|
||||
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)
|
||||
|
||||
|
||||
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]:
|
||||
import huggingface_hub as hf
|
||||
|
||||
|
||||
Reference in New Issue
Block a user