[DCP] Support decode context parallelism on the trtllm_mla decode path (#33926)
Co-authored-by: kpham-sgl <khoa.pham@radixark.ai> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
kpham-sgl
Claude Opus 5
parent
97744189b8
commit
f50b4ad7ae
@@ -10,10 +10,8 @@ returns the rank-local ``(out, lse)`` needed by the cross-rank merge in
|
|||||||
``deepseek_common/attention_forward_methods/forward_mla.py``.
|
``deepseek_common/attention_forward_methods/forward_mla.py``.
|
||||||
|
|
||||||
Non-DCP (``dcp_size == 1``) decode falls through to the base cute-dsl path
|
Non-DCP (``dcp_size == 1``) decode falls through to the base cute-dsl path
|
||||||
unchanged. The DCP metadata helpers below are intentionally duplicated from
|
unchanged. The DCP metadata helpers live on :class:`TRTLLMMLABackend`; this
|
||||||
:mod:`tokenspeed_mla_backend` (they are kernel-agnostic) so that TokenSpeed
|
module only supplies the cute-dsl kernel call and its decode forward.
|
||||||
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
|
from __future__ import annotations
|
||||||
@@ -23,26 +21,18 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
|
|
||||||
import torch
|
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.fixup_zero_kv import fixup_zero_kv_rows
|
||||||
from sglang.kernels.ops.attention.utils import (
|
from sglang.kernels.ops.attention.utils import (
|
||||||
concat_mla_absorb_q_general,
|
concat_mla_absorb_q_general,
|
||||||
mla_quantize_and_rope_for_fp8,
|
mla_quantize_and_rope_for_fp8,
|
||||||
mla_quantize_without_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.environ import envs
|
||||||
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.layers.logits_processor import get_in_autotune_dummy_run
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import is_flashinfer_available
|
from sglang.srt.utils import is_flashinfer_available
|
||||||
|
|
||||||
@@ -51,6 +41,7 @@ if is_flashinfer_available():
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -74,195 +65,6 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
|||||||
backend="cute-dsl",
|
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:
|
|
||||||
metadata = self.forward_decode_metadata
|
|
||||||
if metadata.global_seq_lens_k is None:
|
|
||||||
# Plain decode under DCP also keeps the int32 GLOBAL lens in a
|
|
||||||
# capture-stable buffer (super allocates it only for verify):
|
|
||||||
# the DCP kernel consumes both the rank-local and the global
|
|
||||||
# lens every MLA layer, so both are maintained once per step.
|
|
||||||
metadata.global_seq_lens_k = torch.zeros(
|
|
||||||
(bs,), dtype=torch.int32, device=device
|
|
||||||
)
|
|
||||||
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]
|
|
||||||
# Hoist: refresh the int32 global + rank-local lens once per step
|
|
||||||
# into the capture-stable buffers; forward_decode reads them
|
|
||||||
# instead of recomputing get_dcp_lens + two int32 casts per MLA
|
|
||||||
# layer.
|
|
||||||
metadata.global_seq_lens_k.copy_(seq_lens)
|
|
||||||
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
|
||||||
local_seq_lens = metadata.seq_lens_k
|
|
||||||
|
|
||||||
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
|
|
||||||
)
|
|
||||||
elif (
|
|
||||||
forward_batch.forward_mode.is_decode_or_idle()
|
|
||||||
and self.forward_decode_metadata.seq_lens_k is not None
|
|
||||||
):
|
|
||||||
# Same hoist as verify: the parent stored the int32 GLOBAL
|
|
||||||
# lens in seq_lens_k; keep it as global_seq_lens_k and derive
|
|
||||||
# the rank-local view once per step (forward_decode consumes
|
|
||||||
# both every MLA layer).
|
|
||||||
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.
|
# Kernel + decode forward.
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@@ -289,7 +91,16 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
|||||||
"""
|
"""
|
||||||
if cp_world <= 1:
|
if cp_world <= 1:
|
||||||
return super()._run_decode_kernel(
|
return super()._run_decode_kernel(
|
||||||
query, kv_cache, block_tables, seq_lens, max_seq_len, layer
|
query,
|
||||||
|
kv_cache,
|
||||||
|
block_tables,
|
||||||
|
seq_lens,
|
||||||
|
max_seq_len,
|
||||||
|
layer,
|
||||||
|
causal_seqs=causal_seqs,
|
||||||
|
cp_world=cp_world,
|
||||||
|
cp_rank=cp_rank,
|
||||||
|
return_lse=return_lse,
|
||||||
)
|
)
|
||||||
if causal_seqs is None:
|
if causal_seqs is None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -339,6 +150,8 @@ class CuteDslMLABackend(TRTLLMMLABackend):
|
|||||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
|
if parallel.dcp_enabled and get_in_autotune_dummy_run():
|
||||||
|
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||||
if not parallel.dcp_enabled:
|
if not parallel.dcp_enabled:
|
||||||
return super().forward_decode(
|
return super().forward_decode(
|
||||||
q,
|
q,
|
||||||
|
|||||||
@@ -22,10 +22,11 @@ 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` to share its MLA data preparation and
|
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation, prefill
|
||||||
prefill plumbing. Decode-context parallelism is implemented here because the
|
plumbing, and DCP metadata (rank-local KV lengths and page table). The decode
|
||||||
TokenSpeed decode kernel natively accepts CP rank/world metadata and returns
|
forward lives here because the TokenSpeed decode kernel natively accepts CP
|
||||||
the partial log-sum-exp needed by the cross-rank merge.
|
rank/world metadata and returns the partial log-sum-exp needed by the cross-rank
|
||||||
|
merge.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -34,9 +35,6 @@ 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.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,
|
||||||
@@ -45,16 +43,11 @@ from sglang.kernels.ops.attention.utils import (
|
|||||||
mla_quantize_and_rope_for_fp8,
|
mla_quantize_and_rope_for_fp8,
|
||||||
mla_quantize_without_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.layers.logits_processor import get_in_autotune_dummy_run
|
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -302,165 +295,6 @@ 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,
|
||||||
@@ -517,24 +351,8 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
# FlashInfer autotunes MoE kernels with a synthetic full-model decode
|
|
||||||
# and discards the attention/logits result. On multi-node GB300, the
|
|
||||||
# synthetic full-head DCP metadata can make both the TokenSpeed and
|
|
||||||
# TRTLLM decode kernels surface cudaErrorNvlinkUncorrectable. Skip
|
|
||||||
# attention only inside that explicitly scoped dummy pass. Real
|
|
||||||
# requests and CUDA graph capture continue through TokenSpeed below.
|
|
||||||
if parallel.dcp_enabled and get_in_autotune_dummy_run():
|
if parallel.dcp_enabled and get_in_autotune_dummy_run():
|
||||||
output = torch.zeros(
|
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||||
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
|
|
||||||
dtype=self.q_data_type,
|
|
||||||
device=q.device,
|
|
||||||
)
|
|
||||||
lse = torch.zeros(
|
|
||||||
(q.shape[0], layer.tp_q_head_num),
|
|
||||||
dtype=torch.float32,
|
|
||||||
device=q.device,
|
|
||||||
)
|
|
||||||
return output, lse
|
|
||||||
|
|
||||||
if not parallel.dcp_enabled:
|
if not parallel.dcp_enabled:
|
||||||
return super().forward_decode(
|
return super().forward_decode(
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Optional, Union
|
|||||||
import torch
|
import torch
|
||||||
import triton
|
import triton
|
||||||
|
|
||||||
|
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.fixup_zero_kv import fixup_zero_kv_rows
|
||||||
from sglang.kernels.ops.attention.pad import (
|
from sglang.kernels.ops.attention.pad import (
|
||||||
pad_draft_extend_query as pad_draft_extend_query_triton,
|
pad_draft_extend_query as pad_draft_extend_query_triton,
|
||||||
@@ -50,6 +51,8 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
|
|||||||
FlashInferMLAMultiStepDraftBackend,
|
FlashInferMLAMultiStepDraftBackend,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
||||||
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
|
||||||
is_in_breakable_cuda_graph,
|
is_in_breakable_cuda_graph,
|
||||||
@@ -218,6 +221,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.num_q_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
self.num_q_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
||||||
self.num_kv_heads = config.get_num_kv_heads(get_parallel().attn_tp_size)
|
self.num_kv_heads = config.get_num_kv_heads(get_parallel().attn_tp_size)
|
||||||
self.num_local_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
self.num_local_heads = config.num_attention_heads // get_parallel().attn_tp_size
|
||||||
|
# A DCP decode attends with the query all-gathered across the DCP
|
||||||
|
# group, so the kernel sees attn_dcp_size x this rank's heads. Anything
|
||||||
|
# sized per decode head must use this, not num_q_heads.
|
||||||
|
self.num_decode_q_heads = self.num_q_heads * get_parallel().attn_dcp_size
|
||||||
|
|
||||||
# MLA-specific dimensions
|
# MLA-specific dimensions
|
||||||
self.kv_lora_rank = config.kv_lora_rank
|
self.kv_lora_rank = config.kv_lora_rank
|
||||||
@@ -259,7 +266,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self._multi_ctas_kv_counter_buffer = (
|
self._multi_ctas_kv_counter_buffer = (
|
||||||
make_persistent_multi_ctas_kv_counter_buffer(
|
make_persistent_multi_ctas_kv_counter_buffer(
|
||||||
torch.device(self.device),
|
torch.device(self.device),
|
||||||
self.num_q_heads,
|
self.num_decode_q_heads,
|
||||||
max_batch_size=model_runner.max_running_requests,
|
max_batch_size=model_runner.max_running_requests,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -292,9 +299,15 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
# instead of set_mla_kv_buffer + concat_mla_absorb_q). Disabled under
|
# instead of set_mla_kv_buffer + concat_mla_absorb_q). Disabled under
|
||||||
# async asserts: the fused path writes the pool directly and would
|
# async asserts: the fused path writes the pool directly and would
|
||||||
# skip the pool's OOB probe.
|
# skip the pool's OOB probe.
|
||||||
|
# Also disabled under DCP: unlike its fp8 sibling, set_mla_kv_concat_q
|
||||||
|
# takes no dcp_world_size/dcp_rank, so it writes at the raw virtual
|
||||||
|
# out_cache_loc from every rank, while the reader expects the compacted
|
||||||
|
# row loc // dcp_size written only by the owner. Fall back to the pool's
|
||||||
|
# DCP-aware set_mla_kv_buffer.
|
||||||
self._fused_set_kv_concat_q = (
|
self._fused_set_kv_concat_q = (
|
||||||
self.data_type == torch.bfloat16
|
self.data_type == torch.bfloat16
|
||||||
and not envs.SGLANG_ENABLE_ASYNC_ASSERT.get()
|
and not envs.SGLANG_ENABLE_ASYNC_ASSERT.get()
|
||||||
|
and not get_parallel().dcp_enabled
|
||||||
and can_use_set_mla_kv_concat_q(
|
and can_use_set_mla_kv_concat_q(
|
||||||
self.kv_lora_rank * 2, self.qk_rope_head_dim * 2
|
self.kv_lora_rank * 2, self.qk_rope_head_dim * 2
|
||||||
)
|
)
|
||||||
@@ -333,6 +346,57 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
blocks = triton.cdiv(blocks, constraint_lcm) * constraint_lcm
|
blocks = triton.cdiv(blocks, constraint_lcm) * constraint_lcm
|
||||||
return blocks
|
return blocks
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# DCP metadata (rank-local KV lengths + page table). Kernel-agnostic, so
|
||||||
|
# the whole trtllm_mla family shares it. A no-op when DCP is off.
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
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(
|
def _create_block_kv_indices(
|
||||||
self,
|
self,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
@@ -358,6 +422,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
block_kv_indices,
|
||||||
|
req_pool_indices,
|
||||||
|
self._get_dcp_local_seq_lens(seq_lens),
|
||||||
|
)
|
||||||
|
return block_kv_indices
|
||||||
|
|
||||||
if self.kv_index_translator.is_translating:
|
if self.kv_index_translator.is_translating:
|
||||||
self.kv_index_translator.fill_read_table(
|
self.kv_index_translator.fill_read_table(
|
||||||
out=block_kv_indices,
|
out=block_kv_indices,
|
||||||
@@ -502,6 +574,18 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
metadata.block_kv_indices = block_kv_indices
|
metadata.block_kv_indices = block_kv_indices
|
||||||
metadata.max_seq_len_k = self.max_context_len
|
metadata.max_seq_len_k = self.max_context_len
|
||||||
|
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
if metadata.global_seq_lens_k is None:
|
||||||
|
# A DCP decode consumes both the rank-local and the global
|
||||||
|
# lens, and the branches above allocate this only for verify.
|
||||||
|
metadata.global_seq_lens_k = torch.zeros(
|
||||||
|
(bs,), dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
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)
|
||||||
|
)
|
||||||
|
|
||||||
self.decode_cuda_graph_metadata[bs] = metadata
|
self.decode_cuda_graph_metadata[bs] = metadata
|
||||||
self.forward_decode_metadata = metadata
|
self.forward_decode_metadata = metadata
|
||||||
|
|
||||||
@@ -519,6 +603,11 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
"""
|
"""
|
||||||
metadata = self.decode_cuda_graph_metadata[bs]
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
|
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
return self._apply_dcp_cuda_graph_metadata(
|
||||||
|
bs, req_pool_indices, seq_lens, forward_mode, metadata
|
||||||
|
)
|
||||||
|
|
||||||
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(
|
torch.add(
|
||||||
@@ -565,6 +654,50 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
PAGED_SIZE=self.page_size,
|
PAGED_SIZE=self.page_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _apply_dcp_cuda_graph_metadata(
|
||||||
|
self,
|
||||||
|
bs: int,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
seq_lens: torch.Tensor,
|
||||||
|
forward_mode: ForwardMode,
|
||||||
|
metadata: TRTLLMMLADecodeMetadata,
|
||||||
|
):
|
||||||
|
"""DCP variant of the capture+replay body.
|
||||||
|
|
||||||
|
Refreshes the global and rank-local lengths into the capture-stable
|
||||||
|
buffers once per step, and rebuilds the page table over this rank's
|
||||||
|
cyclic slice.
|
||||||
|
"""
|
||||||
|
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.global_seq_lens_k.copy_(seq_lens)
|
||||||
|
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
||||||
|
local_seq_lens = metadata.seq_lens_k
|
||||||
|
else:
|
||||||
|
seq_lens = seq_lens[:bs]
|
||||||
|
metadata.global_seq_lens_k.copy_(seq_lens)
|
||||||
|
metadata.seq_lens_k.copy_(self._get_dcp_local_seq_lens(seq_lens))
|
||||||
|
local_seq_lens = metadata.seq_lens_k
|
||||||
|
|
||||||
|
self._fill_dcp_block_kv_indices(
|
||||||
|
metadata.block_kv_indices,
|
||||||
|
req_pool_indices[:bs],
|
||||||
|
local_seq_lens,
|
||||||
|
)
|
||||||
|
|
||||||
def get_cuda_graph_seq_len_fill_value(self) -> int:
|
def get_cuda_graph_seq_len_fill_value(self) -> int:
|
||||||
"""Get the fill value for sequence lengths in CUDA graph."""
|
"""Get the fill value for sequence lengths in CUDA graph."""
|
||||||
return 1
|
return 1
|
||||||
@@ -766,6 +899,24 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.forward_decode_metadata.max_seq_len_k = int(max_seq)
|
self.forward_decode_metadata.max_seq_len_k = int(max_seq)
|
||||||
self.forward_decode_metadata.batch_size = bs
|
self.forward_decode_metadata.batch_size = bs
|
||||||
|
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
metadata = self.forward_decode_metadata
|
||||||
|
if (
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_decode_or_idle()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
) and metadata.seq_lens_k is not None:
|
||||||
|
# The branches above stored the global lengths in
|
||||||
|
# seq_lens_k; keep them as global_seq_lens_k and derive the
|
||||||
|
# rank-local view once per step rather than per MLA layer.
|
||||||
|
metadata.global_seq_lens_k = metadata.seq_lens_k
|
||||||
|
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
|
||||||
|
metadata.global_seq_lens_k
|
||||||
|
)
|
||||||
|
metadata.max_seq_len_k = self._get_dcp_local_max_seq_len(
|
||||||
|
metadata.max_seq_len_k
|
||||||
|
)
|
||||||
|
|
||||||
forward_batch.decode_trtllm_mla_metadata = self.forward_decode_metadata
|
forward_batch.decode_trtllm_mla_metadata = self.forward_decode_metadata
|
||||||
else:
|
else:
|
||||||
return super().init_forward_metadata(forward_batch)
|
return super().init_forward_metadata(forward_batch)
|
||||||
@@ -863,15 +1014,29 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
"""Hook for subclasses to swap the decode/spec-verify kernel.
|
"""Hook for subclasses to swap the decode/spec-verify kernel.
|
||||||
|
|
||||||
The DCP arguments belong to the hook contract because forward_extend
|
The DCP arguments belong to the hook contract because forward_extend
|
||||||
passes them on the DCP target-verify path. This implementation does not
|
passes them on the DCP target-verify path.
|
||||||
forward them to the kernel and returns no LSE, so only the DCP-capable
|
|
||||||
subclasses serve them."""
|
The trtllm-gen kernel has no in-kernel DCP support (flashinfer rejects
|
||||||
if cp_world > 1 or return_lse:
|
``enable_dcp=True`` for every backend except ``cute-dsl``), but a
|
||||||
|
``q_len == 1`` decode does not need it: the per-query global causal
|
||||||
|
bound only varies across query rows when ``q_len > 1``, so for a single
|
||||||
|
query token the rank-local page table and ``seq_lens`` already describe
|
||||||
|
the shard completely.
|
||||||
|
|
||||||
|
Every other DCP path is refused. ``q_len > 1`` catches multi-token
|
||||||
|
verify / draft-extend; ``causal_seqs`` and ``return_lse`` catch the
|
||||||
|
single-token verify and draft-extend that ``q_len`` alone lets through
|
||||||
|
(only plain decode requests the LSE the cross-rank merge needs).
|
||||||
|
"""
|
||||||
|
q_len = query.shape[1] if query.dim() == 4 else 1
|
||||||
|
if get_parallel().dcp_enabled and (
|
||||||
|
q_len > 1 or causal_seqs is not None or not return_lse
|
||||||
|
):
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"trtllm_mla does not forward the cyclic DCP metadata to its "
|
"trtllm_mla cannot forward a global causal bound to its decode "
|
||||||
"decode kernel and returns no rank-local LSE for the cross-rank "
|
"kernel, which is required for DCP with q_len > 1 (speculative "
|
||||||
"merge; select cutedsl_mla or tokenspeed_mla for a DCP "
|
"target-verify / draft-extend); select cutedsl_mla or "
|
||||||
"target-verify run"
|
"tokenspeed_mla for a DCP speculative run"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Scale computation for TRTLLM MLA kernel BMM1 operation:
|
# Scale computation for TRTLLM MLA kernel BMM1 operation:
|
||||||
@@ -903,6 +1068,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
seq_lens=seq_lens_i32,
|
seq_lens=seq_lens_i32,
|
||||||
max_seq_len=max_seq_len,
|
max_seq_len=max_seq_len,
|
||||||
bmm1_scale=bmm1_scale,
|
bmm1_scale=bmm1_scale,
|
||||||
|
return_lse=return_lse,
|
||||||
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
|
||||||
**extra_kwargs,
|
**extra_kwargs,
|
||||||
)
|
)
|
||||||
@@ -1045,6 +1211,28 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
dcp_rank=parallel.attn_dcp_rank,
|
dcp_rank=parallel.attn_dcp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _dummy_dcp_decode_for_autotune(
|
||||||
|
self, q: torch.Tensor, layer: RadixAttention
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Skip decode during FlashInfer MoE autotune dummy forwards.
|
||||||
|
|
||||||
|
That pass discards attention/logits. Under DCP the synthetic
|
||||||
|
full-head metadata can overflow the trtllm-gen workspace (and on
|
||||||
|
multi-node GB300 has also produced NVLink errors). Real requests
|
||||||
|
and CUDA-graph capture must not take this path.
|
||||||
|
"""
|
||||||
|
output = torch.zeros(
|
||||||
|
(q.shape[0], layer.tp_q_head_num * layer.v_head_dim),
|
||||||
|
dtype=self.q_data_type,
|
||||||
|
device=q.device,
|
||||||
|
)
|
||||||
|
lse = torch.zeros(
|
||||||
|
(q.shape[0], layer.tp_q_head_num),
|
||||||
|
dtype=torch.float32,
|
||||||
|
device=q.device,
|
||||||
|
)
|
||||||
|
return output, lse
|
||||||
|
|
||||||
def forward_decode(
|
def forward_decode(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor, # q_nope
|
q: torch.Tensor, # q_nope
|
||||||
@@ -1060,6 +1248,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Run forward for decode using TRTLLM MLA kernel."""
|
"""Run forward for decode using TRTLLM MLA kernel."""
|
||||||
|
if get_parallel().dcp_enabled and get_in_autotune_dummy_run():
|
||||||
|
return self._dummy_dcp_decode_for_autotune(q, layer)
|
||||||
|
|
||||||
merge_query = q_rope is not None
|
merge_query = q_rope is not None
|
||||||
fused_fp8_query = None
|
fused_fp8_query = None
|
||||||
if self.data_type == torch.float8_e4m3fn:
|
if self.data_type == torch.float8_e4m3fn:
|
||||||
@@ -1181,6 +1372,11 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.init_forward_metadata(forward_batch)
|
self.init_forward_metadata(forward_batch)
|
||||||
metadata = forward_batch.decode_trtllm_mla_metadata
|
metadata = forward_batch.decode_trtllm_mla_metadata
|
||||||
|
|
||||||
|
if get_parallel().dcp_enabled:
|
||||||
|
return self._forward_decode_dcp(
|
||||||
|
query, kv_cache, metadata, layer, forward_batch
|
||||||
|
)
|
||||||
|
|
||||||
raw_out = self._run_decode_kernel(
|
raw_out = self._run_decode_kernel(
|
||||||
query=query,
|
query=query,
|
||||||
kv_cache=kv_cache,
|
kv_cache=kv_cache,
|
||||||
@@ -1198,6 +1394,48 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
output = raw_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
output = raw_out.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
def _forward_decode_dcp(
|
||||||
|
self,
|
||||||
|
query: torch.Tensor,
|
||||||
|
kv_cache: torch.Tensor,
|
||||||
|
metadata: TRTLLMMLADecodeMetadata,
|
||||||
|
layer: RadixAttention,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
):
|
||||||
|
"""Rank-local MLA decode under DCP, returning ``(out, lse)``.
|
||||||
|
|
||||||
|
The cross-rank merge lives in the model
|
||||||
|
(``deepseek_common/attention_forward_methods/forward_mla.py``), so this
|
||||||
|
returns the rank-local attention state rather than a final output.
|
||||||
|
"""
|
||||||
|
bs = forward_batch.batch_size
|
||||||
|
if metadata.seq_lens_k is not None:
|
||||||
|
local_seq_lens = metadata.seq_lens_k[:bs]
|
||||||
|
else:
|
||||||
|
local_seq_lens = self._get_dcp_local_seq_lens(forward_batch.seq_lens[:bs])
|
||||||
|
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,
|
||||||
|
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)
|
||||||
|
# A rank that owns no slice of a request must contribute a neutral
|
||||||
|
# (out=0, lse=-inf) state, or its garbage rows poison the merge.
|
||||||
|
fixup_zero_kv_rows(
|
||||||
|
output,
|
||||||
|
lse,
|
||||||
|
local_seq_lens,
|
||||||
|
self.q_indptr_decode[: bs + 1],
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
return output.flatten(1), lse
|
||||||
|
|
||||||
def forward_extend(
|
def forward_extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
|
|||||||
@@ -1,87 +0,0 @@
|
|||||||
import unittest
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import patch
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from sglang.srt.layers.attention import tokenspeed_mla_backend as backend_module
|
|
||||||
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
|
||||||
from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLADecodeMetadata
|
|
||||||
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
|
||||||
from sglang.test.test_utils import CustomTestCase
|
|
||||||
|
|
||||||
register_cuda_ci(est_time=60, stage="base-b", runner_config="4-gpu-b200")
|
|
||||||
|
|
||||||
NUM_DRAFT_TOKENS = 8
|
|
||||||
DCP_SIZE = 4
|
|
||||||
DCP_RANK = 2
|
|
||||||
|
|
||||||
|
|
||||||
def _make_backend(bs: int):
|
|
||||||
backend = object.__new__(TokenspeedMLABackend)
|
|
||||||
backend.num_draft_tokens = NUM_DRAFT_TOKENS
|
|
||||||
metadata = TRTLLMMLADecodeMetadata(
|
|
||||||
block_kv_indices=torch.full((bs, 4), -1, dtype=torch.int32, device="cuda"),
|
|
||||||
seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
|
||||||
global_seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
|
||||||
)
|
|
||||||
backend.decode_cuda_graph_metadata = {bs: metadata}
|
|
||||||
return backend, metadata
|
|
||||||
|
|
||||||
|
|
||||||
def _apply(backend, *, bs: int, seq_lens: torch.Tensor, forward_mode):
|
|
||||||
parallel = SimpleNamespace(dcp_enabled=True, dcp_size=DCP_SIZE, dcp_rank=DCP_RANK)
|
|
||||||
with (
|
|
||||||
patch.object(backend_module, "get_parallel", return_value=parallel),
|
|
||||||
patch.object(backend, "_fill_dcp_block_kv_indices") as fill,
|
|
||||||
):
|
|
||||||
backend._apply_cuda_graph_metadata(
|
|
||||||
bs=bs,
|
|
||||||
req_pool_indices=torch.arange(bs, dtype=torch.int32, device="cuda"),
|
|
||||||
seq_lens=seq_lens,
|
|
||||||
forward_mode=forward_mode,
|
|
||||||
)
|
|
||||||
return fill
|
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(torch.cuda.is_available(), "DCP metadata buffers live on CUDA")
|
|
||||||
class TestTokenspeedMLADCPMetadata(CustomTestCase):
|
|
||||||
def test_target_verify_splits_global_and_local_lengths(self):
|
|
||||||
bs = 3
|
|
||||||
backend, metadata = _make_backend(bs)
|
|
||||||
prefix_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
|
||||||
|
|
||||||
fill = _apply(
|
|
||||||
backend,
|
|
||||||
bs=bs,
|
|
||||||
seq_lens=prefix_lens,
|
|
||||||
forward_mode=ForwardMode.TARGET_VERIFY,
|
|
||||||
)
|
|
||||||
|
|
||||||
expected_global = prefix_lens + NUM_DRAFT_TOKENS
|
|
||||||
expected_local = get_dcp_lens(expected_global, DCP_SIZE, DCP_RANK).to(
|
|
||||||
torch.int32
|
|
||||||
)
|
|
||||||
torch.testing.assert_close(metadata.global_seq_lens_k, expected_global)
|
|
||||||
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
|
||||||
fill.assert_called_once()
|
|
||||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
|
||||||
|
|
||||||
def test_decode_does_not_add_draft_tokens(self):
|
|
||||||
bs = 3
|
|
||||||
backend, _ = _make_backend(bs)
|
|
||||||
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
|
||||||
|
|
||||||
fill = _apply(
|
|
||||||
backend, bs=bs, seq_lens=seq_lens, forward_mode=ForwardMode.DECODE
|
|
||||||
)
|
|
||||||
|
|
||||||
expected_local = get_dcp_lens(seq_lens, DCP_SIZE, DCP_RANK).to(torch.int32)
|
|
||||||
fill.assert_called_once()
|
|
||||||
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
unittest.main()
|
|
||||||
@@ -0,0 +1,250 @@
|
|||||||
|
"""DCP cuda-graph metadata for the trtllm_mla backend family.
|
||||||
|
|
||||||
|
The rank-local KV-length and page-table plumbing lives on
|
||||||
|
:class:`TRTLLMMLABackend`, so the same expectations are asserted for the base
|
||||||
|
backend and for both subclasses that inherit it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.attention import trtllm_mla_backend as backend_module
|
||||||
|
from sglang.srt.layers.attention.cutedsl_mla_backend import CuteDslMLABackend
|
||||||
|
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
|
||||||
|
from sglang.srt.layers.attention.trtllm_mla_backend import (
|
||||||
|
TRTLLMMLABackend,
|
||||||
|
TRTLLMMLADecodeMetadata,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dcp.layout import get_dcp_lens
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=60, stage="base-b", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
NUM_DRAFT_TOKENS = 8
|
||||||
|
DCP_SIZE = 4
|
||||||
|
DCP_RANK = 2
|
||||||
|
|
||||||
|
|
||||||
|
def _make_backend(backend_cls, bs: int):
|
||||||
|
backend = object.__new__(backend_cls)
|
||||||
|
backend.num_draft_tokens = NUM_DRAFT_TOKENS
|
||||||
|
metadata = TRTLLMMLADecodeMetadata(
|
||||||
|
block_kv_indices=torch.full((bs, 4), -1, dtype=torch.int32, device="cuda"),
|
||||||
|
seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||||
|
global_seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
|
||||||
|
)
|
||||||
|
backend.decode_cuda_graph_metadata = {bs: metadata}
|
||||||
|
return backend, metadata
|
||||||
|
|
||||||
|
|
||||||
|
def _apply(backend, *, bs: int, seq_lens: torch.Tensor, forward_mode):
|
||||||
|
parallel = SimpleNamespace(dcp_enabled=True, dcp_size=DCP_SIZE, dcp_rank=DCP_RANK)
|
||||||
|
with (
|
||||||
|
patch.object(backend_module, "get_parallel", return_value=parallel),
|
||||||
|
patch.object(backend, "_fill_dcp_block_kv_indices") as fill,
|
||||||
|
):
|
||||||
|
backend._apply_cuda_graph_metadata(
|
||||||
|
bs=bs,
|
||||||
|
req_pool_indices=torch.arange(bs, dtype=torch.int32, device="cuda"),
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
forward_mode=forward_mode,
|
||||||
|
)
|
||||||
|
return fill
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "DCP metadata buffers live on CUDA")
|
||||||
|
class _DCPMetadataTests:
|
||||||
|
backend_cls = None
|
||||||
|
|
||||||
|
def test_target_verify_splits_global_and_local_lengths(self):
|
||||||
|
bs = 3
|
||||||
|
backend, metadata = _make_backend(self.backend_cls, bs)
|
||||||
|
prefix_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||||
|
|
||||||
|
fill = _apply(
|
||||||
|
backend,
|
||||||
|
bs=bs,
|
||||||
|
seq_lens=prefix_lens,
|
||||||
|
forward_mode=ForwardMode.TARGET_VERIFY,
|
||||||
|
)
|
||||||
|
|
||||||
|
expected_global = prefix_lens + NUM_DRAFT_TOKENS
|
||||||
|
expected_local = get_dcp_lens(expected_global, DCP_SIZE, DCP_RANK).to(
|
||||||
|
torch.int32
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(metadata.global_seq_lens_k, expected_global)
|
||||||
|
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
||||||
|
fill.assert_called_once()
|
||||||
|
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||||
|
|
||||||
|
def test_decode_does_not_add_draft_tokens(self):
|
||||||
|
bs = 3
|
||||||
|
backend, metadata = _make_backend(self.backend_cls, bs)
|
||||||
|
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
|
||||||
|
|
||||||
|
fill = _apply(
|
||||||
|
backend, bs=bs, seq_lens=seq_lens, forward_mode=ForwardMode.DECODE
|
||||||
|
)
|
||||||
|
|
||||||
|
expected_local = get_dcp_lens(seq_lens, DCP_SIZE, DCP_RANK).to(torch.int32)
|
||||||
|
fill.assert_called_once()
|
||||||
|
torch.testing.assert_close(fill.call_args.args[2], expected_local)
|
||||||
|
# Plain decode keeps both views in the capture-stable buffers.
|
||||||
|
torch.testing.assert_close(metadata.global_seq_lens_k, seq_lens)
|
||||||
|
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTRTLLMMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||||
|
backend_cls = TRTLLMMLABackend
|
||||||
|
|
||||||
|
|
||||||
|
class TestTokenspeedMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||||
|
backend_cls = TokenspeedMLABackend
|
||||||
|
|
||||||
|
|
||||||
|
class TestCuteDslMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
|
||||||
|
backend_cls = CuteDslMLABackend
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "needs the flashinfer decode hook")
|
||||||
|
class TestTRTLLMMLARejectsDcpMultiTokenQuery(CustomTestCase):
|
||||||
|
"""``_run_decode_kernel`` must refuse every spec path under DCP.
|
||||||
|
|
||||||
|
trtllm-gen takes no global causal bound, which a ``q_len > 1`` batch needs.
|
||||||
|
Single-token spec batches still read as ``q_len == 1``, so the refusal also
|
||||||
|
keys on ``causal_seqs`` and ``return_lse``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _make_backend(self):
|
||||||
|
backend = object.__new__(TRTLLMMLABackend)
|
||||||
|
backend.backend = "trtllm-gen"
|
||||||
|
backend.qk_nope_head_dim = 128
|
||||||
|
backend.kv_lora_rank = 512
|
||||||
|
backend.qk_rope_head_dim = 64
|
||||||
|
backend.workspace_buffer = None
|
||||||
|
backend._multi_ctas_kv_counter_buffer = None
|
||||||
|
return backend
|
||||||
|
|
||||||
|
def _call(
|
||||||
|
self,
|
||||||
|
backend,
|
||||||
|
q_len,
|
||||||
|
*,
|
||||||
|
dcp_enabled,
|
||||||
|
kernel=None,
|
||||||
|
causal_seqs=None,
|
||||||
|
return_lse=False,
|
||||||
|
):
|
||||||
|
parallel = SimpleNamespace(
|
||||||
|
dcp_enabled=dcp_enabled,
|
||||||
|
dcp_size=DCP_SIZE if dcp_enabled else 1,
|
||||||
|
dcp_rank=DCP_RANK if dcp_enabled else 0,
|
||||||
|
)
|
||||||
|
flashinfer_stub = SimpleNamespace(
|
||||||
|
decode=SimpleNamespace(
|
||||||
|
trtllm_batch_decode_with_kv_cache_mla=kernel or (lambda **kw: None)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
with (
|
||||||
|
patch.object(backend_module, "get_parallel", return_value=parallel),
|
||||||
|
patch.object(backend_module, "flashinfer", flashinfer_stub),
|
||||||
|
patch.object(backend, "_compute_decode_bmm1_scale", return_value=1.0),
|
||||||
|
):
|
||||||
|
return backend._run_decode_kernel(
|
||||||
|
query=torch.zeros(
|
||||||
|
(2, q_len, 16, 576), dtype=torch.bfloat16, device="cuda"
|
||||||
|
),
|
||||||
|
kv_cache=torch.zeros((4, 1, 64, 576), dtype=torch.bfloat16),
|
||||||
|
block_tables=torch.zeros((2, 4), dtype=torch.int32, device="cuda"),
|
||||||
|
seq_lens=torch.ones(2, dtype=torch.int32, device="cuda"),
|
||||||
|
max_seq_len=64,
|
||||||
|
layer=SimpleNamespace(scaling=1.0, k_scale_float=None),
|
||||||
|
causal_seqs=causal_seqs,
|
||||||
|
return_lse=return_lse,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multi_token_query_under_dcp_raises(self):
|
||||||
|
with self.assertRaises(NotImplementedError):
|
||||||
|
self._call(self._make_backend(), q_len=NUM_DRAFT_TOKENS, dcp_enabled=True)
|
||||||
|
|
||||||
|
def test_explicit_causal_bound_under_dcp_raises_at_q_len_one(self):
|
||||||
|
# Nothing forces speculative_num_draft_tokens > 1, so a q_len == 1
|
||||||
|
# target-verify is expressible and must not slip past the q_len proxy.
|
||||||
|
with self.assertRaises(NotImplementedError):
|
||||||
|
self._call(
|
||||||
|
self._make_backend(),
|
||||||
|
q_len=1,
|
||||||
|
dcp_enabled=True,
|
||||||
|
causal_seqs=torch.ones(2, dtype=torch.int32, device="cuda"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_single_token_query_under_dcp_is_allowed(self):
|
||||||
|
# The premise of trtllm_mla DCP decode: q_len == 1 needs no global
|
||||||
|
# bound, so the guard must not swallow the path it exists to protect.
|
||||||
|
calls = []
|
||||||
|
self._call(
|
||||||
|
self._make_backend(),
|
||||||
|
q_len=1,
|
||||||
|
dcp_enabled=True,
|
||||||
|
kernel=lambda **kw: calls.append(kw),
|
||||||
|
return_lse=True,
|
||||||
|
)
|
||||||
|
self.assertEqual(len(calls), 1)
|
||||||
|
|
||||||
|
def test_single_token_draft_extend_under_dcp_raises(self):
|
||||||
|
# A single-token draft-extend reads as q_len == 1 with no causal_seqs;
|
||||||
|
# skipping the cross-rank merge (no LSE requested) is the only signal.
|
||||||
|
with self.assertRaises(NotImplementedError):
|
||||||
|
self._call(self._make_backend(), q_len=1, dcp_enabled=True)
|
||||||
|
|
||||||
|
def test_multi_token_query_without_dcp_is_allowed(self):
|
||||||
|
calls = []
|
||||||
|
self._call(
|
||||||
|
self._make_backend(),
|
||||||
|
q_len=NUM_DRAFT_TOKENS,
|
||||||
|
dcp_enabled=False,
|
||||||
|
kernel=lambda **kw: calls.append(kw),
|
||||||
|
)
|
||||||
|
self.assertEqual(len(calls), 1)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDcpDecodeLayout(CustomTestCase):
|
||||||
|
"""Rank-local length math the decode page table above is built from."""
|
||||||
|
|
||||||
|
SIZES = [1, 2, 3, 4, 8]
|
||||||
|
LENS = list(range(0, 41))
|
||||||
|
|
||||||
|
def test_ranks_partition_the_global_length(self):
|
||||||
|
lens = torch.tensor(self.LENS, dtype=torch.int32)
|
||||||
|
for n in self.SIZES:
|
||||||
|
total = sum(
|
||||||
|
get_dcp_lens(lens, n, rank).to(torch.int64) for rank in range(n)
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(total, lens.to(torch.int64)),
|
||||||
|
f"per-rank lengths do not sum to the global length at n={n}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_newest_token_is_owned_by_exactly_one_rank(self):
|
||||||
|
# A decode step appends one token; the cross-rank merge double counts
|
||||||
|
# or drops it unless exactly one rank sees its length grow.
|
||||||
|
for n in self.SIZES:
|
||||||
|
for global_len in self.LENS[1:]:
|
||||||
|
prev = torch.tensor([global_len - 1], dtype=torch.int32)
|
||||||
|
cur = torch.tensor([global_len], dtype=torch.int32)
|
||||||
|
grew = [
|
||||||
|
int(get_dcp_lens(cur, n, rank).item())
|
||||||
|
- int(get_dcp_lens(prev, n, rank).item())
|
||||||
|
for rank in range(n)
|
||||||
|
]
|
||||||
|
self.assertEqual(sum(grew), 1, f"n={n}, global_len={global_len}")
|
||||||
|
self.assertEqual(grew[(global_len - 1) % n], 1)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user