Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 13 additions & 3 deletions python/sglang/srt/layers/attention/trtllm_mla_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -1240,12 +1240,14 @@ def _set_kv_and_concat_q_fp8_fused(
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.
"""Skip DCP decode / target-verify during FlashInfer 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.
multi-node GB300 has also produced NVLink errors), and the FlashInfer
kernels (trtllm-gen, cute-dsl) start their own tuning, whose synthetic
inputs can OOM on some ranks only and hang the cross-rank reduction.
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),
Expand Down Expand Up @@ -1476,6 +1478,14 @@ def forward_extend(
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# A speculative runner's autotune dummy forward is TARGET_VERIFY-shaped,
# so it never reaches forward_decode's guard.
if (
forward_batch.forward_mode.is_target_verify()
and get_parallel().dcp_enabled
and get_in_autotune_dummy_run()
):
return self._dummy_dcp_decode_for_autotune(q, layer)

# The fallback belongs to genuine extend forwards only. Target-verify /
# draft-extend must never honor it: `forward_prefill_metadata` is a
Expand Down
63 changes: 63 additions & 0 deletions test/registered/dcp/test_trtllm_mla_family_dcp_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
backend and for both subclasses that inherit it.
"""

import contextlib
import unittest
from types import SimpleNamespace
from unittest.mock import patch
Expand All @@ -19,6 +20,7 @@
TRTLLMMLADecodeMetadata,
)
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
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
Expand All @@ -30,6 +32,10 @@
DCP_RANK = 2


class _RealVerifyPath(Exception):
"""Raised by the stubs the real verify path reaches first."""


def _make_backend(backend_cls, bs: int):
backend = object.__new__(backend_cls)
backend.num_draft_tokens = NUM_DRAFT_TOKENS
Expand Down Expand Up @@ -98,6 +104,63 @@ def test_decode_does_not_add_draft_tokens(self):
torch.testing.assert_close(metadata.global_seq_lens_k, seq_lens)
torch.testing.assert_close(metadata.seq_lens_k, expected_local)

def _verify_extend(self, *, dcp_enabled: bool, in_autotune: bool):
heads, v_head_dim, n = 4, 512, 3 * NUM_DRAFT_TOKENS
backend = object.__new__(self.backend_cls)
backend.data_type = backend.q_data_type = torch.bfloat16
backend._decode_kernel_loc = None

def real_path(*args, **kwargs):
raise _RealVerifyPath

backend.token_to_kv_pool = SimpleNamespace(set_mla_kv_buffer=real_path)
backend._run_decode_kernel = real_path
layer = SimpleNamespace(tp_q_head_num=heads, v_head_dim=v_head_dim)
forward_batch = SimpleNamespace(
forward_mode=ForwardMode.TARGET_VERIFY, out_cache_loc=None
)
k = torch.zeros((n, 1, v_head_dim), dtype=torch.bfloat16, device="cuda")
k_rope = torch.zeros((n, 1, 64), dtype=torch.bfloat16, device="cuda")
parallel = SimpleNamespace(dcp_enabled=dcp_enabled)
autotune = (
autotune_dummy_run_mode(run_lm_head=False)
if in_autotune
else contextlib.nullcontext()
)
with (
patch.object(backend_module, "get_parallel", return_value=parallel),
autotune,
):
return backend.forward_extend(
torch.zeros(
(n, heads, v_head_dim), dtype=torch.bfloat16, device="cuda"
),
k,
None,
layer,
forward_batch,
save_kv_cache=True,
q_rope=torch.zeros((n, heads, 64), dtype=torch.bfloat16, device="cuda"),
k_rope=k_rope,
)

def test_autotune_verify_under_dcp_skips_the_kernel(self):
"""A speculative autotune dummy verify under DCP must not run the kernel;
a FlashInfer kernel would start its own tuning, which can hang ranks."""
out, lse = self._verify_extend(dcp_enabled=True, in_autotune=True)
n = 3 * NUM_DRAFT_TOKENS
self.assertEqual((out.shape, out.dtype), ((n, 4 * 512), torch.bfloat16))
self.assertEqual((lse.shape, lse.dtype), ((n, 4), torch.float32))
self.assertFalse(out.any() or lse.any())

def test_real_verify_under_dcp_takes_the_real_path(self):
with self.assertRaises(_RealVerifyPath):
self._verify_extend(dcp_enabled=True, in_autotune=False)

def test_autotune_verify_without_dcp_takes_the_real_path(self):
with self.assertRaises(_RealVerifyPath):
self._verify_extend(dcp_enabled=False, in_autotune=True)


class TestTRTLLMMLADCPMetadata(_DCPMetadataTests, CustomTestCase):
backend_cls = TRTLLMMLABackend
Expand Down
Loading