Deprecate record_nolora_graph dual MoE CUDA graph capture (#24314)

This commit is contained in:
Sam Shleifer
2026-05-15 19:20:58 -07:00
committed by GitHub
parent b674007026
commit ce2506e1c6
4 changed files with 39 additions and 129 deletions
-25
View File
@@ -152,7 +152,6 @@ MOE_A2A_BACKEND: Optional[MoeA2ABackend] = None
MOE_RUNNER_BACKEND: Optional[MoeRunnerBackend] = None
SPECULATIVE_MOE_RUNNER_BACKEND: Optional[MoeRunnerBackend] = None
SPECULATIVE_MOE_A2A_BACKEND: Optional[MoeA2ABackend] = None
RECORD_NOLORA_GRAPH: bool = False
DEEPEP_MODE: Optional[DeepEPMode] = None
IS_TBO_ENABLED: Optional[bool] = None
IS_SBO_ENABLED: Optional[bool] = None
@@ -167,7 +166,6 @@ def initialize_moe_config(server_args: ServerArgs):
global MOE_RUNNER_BACKEND
global SPECULATIVE_MOE_RUNNER_BACKEND
global SPECULATIVE_MOE_A2A_BACKEND
global RECORD_NOLORA_GRAPH
global DEEPEP_MODE
global DEEPEP_CONFIG
global IS_TBO_ENABLED
@@ -178,25 +176,6 @@ def initialize_moe_config(server_args: ServerArgs):
MOE_A2A_BACKEND = MoeA2ABackend(server_args.moe_a2a_backend)
MOE_RUNNER_BACKEND = MoeRunnerBackend(server_args.moe_runner_backend)
# Dual CUDA graphs only validated for triton MoE backends.
_triton_ok = MOE_RUNNER_BACKEND in (
MoeRunnerBackend.TRITON,
MoeRunnerBackend.TRITON_KERNELS,
)
if (
bool(server_args.record_nolora_graph)
and bool(server_args.enable_lora)
and not _triton_ok
):
logger.warning(
f"record_nolora_graph only validated for triton MoE backend, "
f"but moe_runner_backend={server_args.moe_runner_backend}. Disabling."
)
RECORD_NOLORA_GRAPH = (
bool(server_args.record_nolora_graph)
and bool(server_args.enable_lora)
and _triton_ok
)
SPECULATIVE_MOE_RUNNER_BACKEND = (
MoeRunnerBackend(server_args.speculative_moe_runner_backend)
if server_args.speculative_moe_runner_backend is not None
@@ -252,10 +231,6 @@ def get_speculative_moe_a2a_backend() -> MoeA2ABackend:
return SPECULATIVE_MOE_A2A_BACKEND
def should_record_nolora_graph() -> bool:
return RECORD_NOLORA_GRAPH
def get_deepep_mode() -> DeepEPMode:
global DEEPEP_MODE
if DEEPEP_MODE is None:
+17 -19
View File
@@ -336,13 +336,23 @@ def _add_lora_gate_up_delta(
)
if get_is_capture_mode():
from sglang.srt.model_executor.cuda_graph_runner import get_capture_lora_variant
# Record LoRA kernels for lora graph; skip for nolora graph.
if get_capture_lora_variant() == "nolora":
return
if lora_info is None or lora_info.max_lora_rank == 0:
# During CUDA graph capture, always enter the LoRA path so that
# the LoRA kernels are recorded in the graph. adapter_enabled is
# all-zeros during capture, so the Triton kernel early-exits per
# program (zero overhead). During replay the tensor is updated
# in-place with the real adapter mask before graph.replay().
has_active_lora = True
else:
num_loras = len(lora_info.lora_ranks)
has_active_lora = (
(
lora_info.adapter_enabled[:num_loras]
* (lora_info.lora_ranks > 0).to(lora_info.adapter_enabled.dtype)
)
.any()
.item()
)
if not has_active_lora or lora_info is None or lora_info.max_lora_rank == 0:
return
M, top_k, gate_up_dim = intermediate_cache.shape
@@ -436,12 +446,6 @@ def _add_lora_down_delta(
if lora_info.max_lora_rank == 0:
return
if get_is_capture_mode():
from sglang.srt.model_executor.cuda_graph_runner import get_capture_lora_variant
if get_capture_lora_variant() == "nolora":
return
M, top_k, hidden_dim = intermediate_cache.shape
down_lora_a = lora_info.down_lora_a_weights
@@ -517,12 +521,6 @@ def build_lora_hooks(
if lora_info is None or lora_info.max_lora_rank == 0:
return LoRAHooks()
if get_is_capture_mode():
from sglang.srt.model_executor.cuda_graph_runner import get_capture_lora_variant
if get_capture_lora_variant() == "nolora":
return LoRAHooks()
# Compute alignment / mapping (once, shared by both hooks)
token_lora_mapping: torch.Tensor | None = None
sorted_token_ids_reshaped: torch.Tensor | None = None
@@ -54,11 +54,7 @@ from sglang.srt.layers.dp_attention import (
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
from sglang.srt.layers.moe.utils import (
get_deepep_mode,
get_moe_a2a_backend,
should_record_nolora_graph,
)
from sglang.srt.layers.moe.utils import get_deepep_mode, get_moe_a2a_backend
from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
@@ -391,9 +387,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
# Detect whether the current forward pass is in capture mode
is_capture_mode = False
# When capturing dual MoE backends, tracks which variant is being captured.
# None = not dual, "lora" = capturing lora variant, "nolora" = capturing nolora variant.
_capture_lora_variant: Optional[str] = None
def get_is_capture_mode():
@@ -406,16 +399,6 @@ def compile_in_capture_mode(func):
return func
def get_capture_lora_variant() -> Optional[str]:
"""Return the lora variant being captured, or None if not in dual capture."""
return _capture_lora_variant
def _set_capture_lora_variant(variant: Optional[str]):
global _capture_lora_variant
_capture_lora_variant = variant
@contextmanager
def model_capture_mode():
global is_capture_mode
@@ -555,19 +538,6 @@ def set_global_graph_memory_pool(val):
global_graph_memory_pool = val
def _default_make_graph_key(bs, stream_idx=None, variant_label=None):
"""Build a graph dict key from batch size, stream index, and lora variant.
Standalone function so it can be used by CudaGraphRunner.capture() even when
called on subclasses (e.g. EAGLEDraftCudaGraphRunner) that don't inherit from
CudaGraphRunner and thus lack the method.
"""
key = bs if stream_idx is None else f"{stream_idx}_{bs}"
if variant_label is not None:
key = f"{variant_label}_{key}"
return key
class CudaGraphRunner:
"""A CudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile."""
@@ -614,7 +584,6 @@ class CudaGraphRunner:
self.nsa_enable_prefill_cp = is_nsa_enable_prefill_cp()
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
self.record_nolora_graph = should_record_nolora_graph()
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
self.is_dllm = self.dllm_config is not None
@@ -740,20 +709,6 @@ class CudaGraphRunner:
def _cache_loc_dtype(self):
return torch.int64
def _make_graph_key(self, bs, stream_idx=None, variant_label=None):
"""Build a graph dict key from batch size, stream index, and lora variant."""
return _default_make_graph_key(bs, stream_idx, variant_label)
def _resolve_lora_variant(self, forward_batch: ForwardBatch):
"""Return the variant label for the given batch, or None if dual backends are off."""
if not getattr(self, "record_nolora_graph", False):
return None
if forward_batch.lora_ids is not None and any(
uid is not None for uid in forward_batch.lora_ids
):
return "lora"
return "nolora"
def can_run(self, forward_batch: ForwardBatch):
# Disable for token embedding overrides (dynamic per-request)
if forward_batch.replace_embeds is not None:
@@ -769,9 +724,9 @@ class CudaGraphRunner:
else:
cuda_graph_bs = forward_batch.batch_size
variant_label = self._resolve_lora_variant(forward_batch)
stream_idx = get_current_stream_idx() if self.enable_pdmux else None
graph_key = self._make_graph_key(cuda_graph_bs, stream_idx, variant_label)
graph_key = cuda_graph_bs
if self.enable_pdmux:
graph_key = f"{get_current_stream_idx()}_{cuda_graph_bs}"
is_bs_supported = (
graph_key in self.graphs
@@ -866,13 +821,6 @@ class CudaGraphRunner:
if get_tensor_model_parallel_rank() == 0
else reversed(self.capture_bs)
)
# When record_nolora_graph is set, capture each batch size twice:
# once with LoRA hooks and once without.
lora_variants = (
[("lora", True), ("nolora", False)]
if getattr(self, "record_nolora_graph", False)
else [(None, None)]
)
for i, bs in enumerate(capture_range):
if get_tensor_model_parallel_rank() == 0:
avail_mem = get_available_gpu_memory(
@@ -884,21 +832,20 @@ class CudaGraphRunner:
f"Capturing batches ({bs=} {avail_mem=:.2f} GB)"
)
for variant_label, variant_has_lora in lora_variants:
_set_capture_lora_variant(variant_label)
with patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs,
tp_group=self.model_runner.tp_group,
) as forward:
(
graph,
output_buffers,
) = self.capture_one_batch_size(bs, forward, stream_idx)
key = _default_make_graph_key(bs, stream_idx, variant_label)
self.graphs[key] = graph
self.output_buffers[key] = output_buffers
with patch_model(
self.model_runner.model,
bs in self.compile_bs,
num_tokens=bs * self.num_tokens_per_bs,
tp_group=self.model_runner.tp_group,
) as forward:
(
graph,
output_buffers,
) = self.capture_one_batch_size(bs, forward, stream_idx)
# For pd_multiplexing, we need to save the graph and output buffers
key = bs if stream_idx is None else f"{stream_idx}_{bs}"
self.graphs[key] = graph
self.output_buffers[key] = output_buffers
# Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes
@@ -918,8 +865,6 @@ class CudaGraphRunner:
self.stream = graph_capture_context.stream
_capture_one_stream(i)
_set_capture_lora_variant(None)
if self.enable_profile_cuda_graph:
self._post_process_after_profile(prof)
@@ -1330,9 +1275,10 @@ class CudaGraphRunner:
)
# Replay
variant_label = self._resolve_lora_variant(forward_batch)
stream_idx = get_current_stream_idx() if self.enable_pdmux else None
graph_key = self._make_graph_key(self.bs, stream_idx, variant_label)
if self.enable_pdmux:
graph_key = f"{get_current_stream_idx()}_{self.bs}"
else:
graph_key = self.bs
ctx = (
self.model_runner.device_timer.wrap(
metadata={
-9
View File
@@ -613,7 +613,6 @@ class ServerArgs:
"none", "deepep", "mooncake", "nixl", "mori", "ascend_fuseep", "flashinfer"
] = "none"
moe_runner_backend: str = "auto"
record_nolora_graph: bool = True
flashinfer_mxfp4_moe_precision: Literal["default", "bf16"] = "default"
enable_flashinfer_allreduce_fusion: bool = False
enforce_disable_flashinfer_allreduce_fusion: bool = False
@@ -5977,14 +5976,6 @@ class ServerArgs:
default=ServerArgs.moe_runner_backend,
help="Choose the runner backend for MoE.",
)
parser.add_argument(
"--record-nolora-graph",
action=argparse.BooleanOptionalAction,
default=ServerArgs.record_nolora_graph,
help="Capture a second set of CUDA graphs without LoRA hooks. "
"Batches without active adapters replay the faster nolora graph. "
"Enabled by default.",
)
parser.add_argument(
"--flashinfer-mxfp4-moe-precision",
type=str,