[BCG][5/N] MLA Fully Support (#33661)

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Yuwei An
2026-08-09 23:52:14 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent aea78d1e73
commit 06f32bab6b
7 changed files with 42 additions and 78 deletions
@@ -708,6 +708,10 @@ def _flash_attn_fwd(
else:
num_splits = 1
if qv is not None:
# The qv kernel has no split-KV variant.
num_splits = 1
is_split_kv = num_splits > 1
if is_split_kv:
out_partial = torch.empty(
-20
View File
@@ -507,9 +507,6 @@ class ModelConfig:
self.is_multimodal_breakable_cuda_graph_supported = enable_multimodal and (
is_multimodal_breakable_cuda_graph_supported(self.hf_config.architectures)
)
self.is_mla_breakable_cuda_graph_supported = (
is_mla_breakable_cuda_graph_supported(self.hf_config.architectures)
)
self.dtype = _get_and_verify_dtype(self.hf_text_config, dtype)
# Derive context length and model shapes
@@ -1898,15 +1895,6 @@ multimodal_breakable_cuda_graph_supported_model_archs = [
"Qwen3_5MoeForConditionalGeneration",
]
# MLA archs validated to run breakable CUDA graph when it is explicitly
# requested (--cuda-graph-backend-prefill=breakable bypasses the ServerArgs
# disable rules). Dispatch pins the absorbed MLA path inside capture/replay
# for these archs, so the prefill runner's MHA-companion prefix restrictions
# do not apply (see PrefillCudaGraphRunner.mla_pinned_under_bcg).
mla_breakable_cuda_graph_supported_model_archs = [
"KimiK3ForConditionalGeneration",
]
if external_mm_model_arch := envs.SGLANG_EXTERNAL_MM_MODEL_ARCH.get():
multimodal_model_archs.append(external_mm_model_arch)
@@ -1981,14 +1969,6 @@ def is_multimodal_breakable_cuda_graph_supported(model_architectures: List[str])
)
def is_mla_breakable_cuda_graph_supported(model_architectures: List[str]):
"""Whether an MLA arch may keep prefill breakable CUDA graph enabled."""
return any(
arch in mla_breakable_cuda_graph_supported_model_archs
for arch in model_architectures
)
# SequenceClassification models that use CrossEncodingPooler
_cross_encoding_pooler_archs = [
"BertForSequenceClassification",
@@ -52,6 +52,9 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
from sglang.srt.layers.attention.unified_mem_hooks import unified_mla_hooks
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
is_in_breakable_cuda_graph,
)
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph,
)
@@ -566,8 +569,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
) -> None:
has_prefix = any(forward_batch.extend_prefix_lens_cpu)
fallback_to_flashinfer_impl = (
self.disable_chunked_prefix_cache and has_prefix
) or is_in_tc_piecewise_cuda_graph()
(self.disable_chunked_prefix_cache and has_prefix)
or is_in_tc_piecewise_cuda_graph()
or is_in_breakable_cuda_graph()
)
if fallback_to_flashinfer_impl:
super().init_mha_chunk_metadata(
forward_batch, disable_flashinfer_ragged=True
@@ -647,11 +652,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
):
# For extend batch with prefix length > 0, fallback to ragged kernel implemented in flashinfer MLA backend
# when chunked prefix cache is disabled.
# Also fallback to flashinfer MLA backend when in piecewise cuda graph, since it only supports MLA forward mode.
# Also fallback to flashinfer MLA backend under a captured prefill graph
has_prefix = any(forward_batch.extend_prefix_lens_cpu)
fallback_to_flashinfer_impl = (
self.disable_chunked_prefix_cache and has_prefix
) or is_in_tc_piecewise_cuda_graph()
(self.disable_chunked_prefix_cache and has_prefix)
or is_in_tc_piecewise_cuda_graph()
or is_in_breakable_cuda_graph()
)
if fallback_to_flashinfer_impl:
super().init_forward_metadata(forward_batch)
@@ -50,7 +50,6 @@ import tqdm
from sglang.kernels.ops.kvcache.kv_indices import (
create_chunked_prefix_cache_kv_indices,
)
from sglang.srt.configs.model_config import is_deepseek_dsa
from sglang.srt.distributed.parallel_state import graph_capture
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
from sglang.srt.layers.cp.bcg import (
@@ -119,6 +118,7 @@ from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.utils import (
get_available_gpu_memory,
is_cuda,
is_npu,
require_attn_tp_gather,
require_gathered_buffer,
@@ -249,11 +249,6 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
buffer population, attention metadata init, and output slicing.
"""
# DSA forces use_mha=False in BCG capture/replay, so the sparse path
# serves any prefix and the MHA-prefix ban does not apply. Class
# default keeps __new__-built test instances on the ban.
dsa_sparse_prefill_forced: bool = False
def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner)
# --- model flags ----------------------------------------------
@@ -330,21 +325,11 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
source=self.buffers,
)
self.dsa_sparse_prefill_forced = is_deepseek_dsa(
self.model_runner.model_config.hf_config
)
self.attention_layers = self.model_runner.attention_layers
self.mha_companion_layers = self.model_runner.mha_companion_layers
self.has_mha_companion_layers = any(
layer is not None for layer in self.mha_companion_layers
)
# Archs on the MLA-BCG allowlist pin the absorbed MLA path inside
# capture/replay (attention_backend_handler), so the MHA companion is
# never captured and the MHA-prefix restrictions below don't apply.
self.mla_pinned_under_bcg = (
self.model_runner.model_config.is_mla_breakable_cuda_graph_supported
)
self.moe_layers = self.model_runner.moe_layers
self.moe_fusions = self.model_runner.moe_fusions
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
@@ -1058,16 +1043,11 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
return False
if replace_embeds is not None:
return False
# A prefix forces the MHA companion path, whose captured state is
# frozen prefix-free; DSA models are exempt (capture/replay force
# the sparse path, which takes any prefix via device metadata), as
# are archs on the MLA-BCG allowlist (they pin the absorbed MLA path
# inside capture/replay, so the MHA companion is never captured).
# Off CUDA, BCG takes the MHA companion, whose prefix path is uncapturable.
if (
self.prefill_backend_name == Backend.BREAKABLE
and self.has_mha_companion_layers
and not self.dsa_sparse_prefill_forced
and not self.mla_pinned_under_bcg
and not is_cuda()
and prefix_lens is not None
and any(prefix_lens)
):
@@ -1577,7 +1557,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
if (
isinstance(self.backend, BreakableCudaGraphBackend)
and self.has_mha_companion_layers
and not self.mla_pinned_under_bcg
and not is_cuda()
):
self._restore_mha_capture_state(static_forward_batch)
+3 -18
View File
@@ -4382,11 +4382,8 @@ class ServerArgs:
if (
self.cuda_graph_config.prefill.backend == Backend.BREAKABLE
and self.get_model_config().is_multimodal_piecewise_cuda_graph_supported
# Keep trtllm_mla on the preferred breakable path. Its current
# breakable compatibility rule disables the graph, avoiding the
# tc_piecewise FlashInfer paged-MLA fallback; once breakable gains
# native support, that rule can be relaxed without re-enabling the
# deprecated tc_piecewise path.
# Keep trtllm_mla on the preferred breakable path, which now serves
# MLA by falling back to the flashinfer MLA impl for extend.
and self._resolved_attention_backends()[0] != "trtllm_mla"
):
logger.info(
@@ -4490,22 +4487,10 @@ class ServerArgs:
memory-saver rejection in its own __init__; config-time rules can be
added here as they're discovered.
"""
from sglang.srt.configs.model_config import (
is_deepseek_dsa,
is_deepseek_v4,
is_nemotron_h,
)
from sglang.srt.configs.model_config import is_deepseek_v4, is_nemotron_h
from sglang.srt.layers.cp.bcg import supports_prefill_cp_bcg
rules = [
# MLA prefill under BCG takes forward_mha, which has no eager
# breaks. DSA is exempt: BCG forces the sparse path, whose
# indexer already splits eagerly.
(
"MLA attention (non-DSA)",
lambda: self.use_mla_backend()
and not is_deepseek_dsa(self.get_model_config().hf_config),
),
# NemotronH's hybrid Mamba2 prefill is not BCG-safe: the mamba
# state-track write is not wired into the captured buffers, so a
# replay can commit a cache slot it never wrote.
+7 -4
View File
@@ -3527,14 +3527,17 @@ def dispose_tensor(x: torch.Tensor):
interfering with torch.compile's memory tracking and graph recording.
"""
# Skip disposal during piecewise CUDA graph capture/replay: freeing the
# backing storage would invalidate addresses recorded in the graph.
# Local import avoids a circular dependency.
# Skip disposal under a captured prefill graph (piecewise or breakable):
# freeing the backing storage would invalidate addresses recorded in the
# graph. Local imports avoid a circular dependency.
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import (
is_in_breakable_cuda_graph,
)
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
is_in_tc_piecewise_cuda_graph,
)
if is_in_tc_piecewise_cuda_graph():
if is_in_tc_piecewise_cuda_graph() or is_in_breakable_cuda_graph():
return
from sglang.srt.runtime_context import get_flags
@@ -36,7 +36,6 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
runner._capture_chunked_prefix = False
runner.prefill_backend_name = backend
runner.has_mha_companion_layers = backend == Backend.BREAKABLE
runner.mla_pinned_under_bcg = False
runner.capture_hidden_mode = CaptureHiddenMode.NULL
runner.capture_num_tokens = [4, 16]
runner.max_num_tokens = 16
@@ -93,12 +92,14 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.TC_PIECEWISE)
disable_if_incompatible.assert_called_once()
def test_trtllm_mla_stays_on_breakable_and_is_disabled_by_compatibility(self):
def test_trtllm_mla_stays_on_breakable(self):
args = ServerArgs(model_path="dummy")
# The MLA rule reads hf_config to exempt DSA models, so the stub needs
# an architecture that is MLA but not DSA.
# trtllm_mla skips the tc_piecewise upgrade and keeps breakable, which
# now serves MLA by falling back to the flashinfer MLA impl for extend.
args.model_config = SimpleNamespace(
is_multimodal_piecewise_cuda_graph_supported=True,
is_multimodal=False,
is_multimodal_breakable_cuda_graph_supported=False,
hf_config=SimpleNamespace(architectures=["DeepseekV2ForCausalLM"]),
)
args.cuda_graph_config = CudaGraphConfig(
@@ -116,7 +117,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
):
args._apply_cuda_graph_compatibility()
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.DISABLED)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
def test_explicit_tc_piecewise_overrides_trtllm_mla_default(self):
args = ServerArgs(model_path="dummy")
@@ -144,12 +145,16 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
self.assertTrue(runner.can_run_graph(self._make_multimodal_forward_batch()))
def test_breakable_prefill_rejects_nonzero_prefix(self):
def test_breakable_prefill_takes_nonzero_prefix_on_cuda_only(self):
runner = self._make_prefill_runner(Backend.BREAKABLE)
forward_batch = self._make_multimodal_forward_batch()
forward_batch.extend_prefix_lens_cpu = [1]
self.assertFalse(runner.can_run_graph(forward_batch))
target = "sglang.srt.model_executor.runner.prefill_cuda_graph_runner.is_cuda"
with patch(target, return_value=True):
self.assertTrue(runner.can_run_graph(forward_batch))
with patch(target, return_value=False):
self.assertFalse(runner.can_run_graph(forward_batch))
def test_embedding_gemma_forces_breakable_prefill(self):
args = ServerArgs(model_path="dummy")