Deprecate record_nolora_graph dual MoE CUDA graph capture (#24314)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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={
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user