From e4bf0043fe802222fb027c1b6e6f914ee7b0cbb5 Mon Sep 17 00:00:00 2001 From: Jiajun Li Date: Wed, 24 Jun 2026 16:56:07 -0700 Subject: [PATCH] [fix] Skip routed expert capture for draft model under spec v2 (#26980) Co-authored-by: Claude Opus 4.8 (1M context) --- .../srt/hardware_backend/npu/moe/topk.py | 13 ++++---- .../srt/layers/moe/fused_moe_triton/layer.py | 3 ++ python/sglang/srt/layers/moe/topk.py | 30 +++++++++++++++---- .../sglang/srt/model_executor/model_runner.py | 17 ++++++++++- .../srt/state_capturer/routed_experts.py | 18 ++++++++++- 5 files changed, 67 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/moe/topk.py b/python/sglang/srt/hardware_backend/npu/moe/topk.py index 333fefadb..887c25a21 100644 --- a/python/sglang/srt/hardware_backend/npu/moe/topk.py +++ b/python/sglang/srt/hardware_backend/npu/moe/topk.py @@ -5,8 +5,11 @@ from sgl_kernel_npu.norm.l1_norm import l1_norm from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location_dispatch import topk_ids_logical_to_physical -from sglang.srt.layers.moe.topk import StandardTopKOutput, select_experts -from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer +from sglang.srt.layers.moe.topk import ( + StandardTopKOutput, + capture_routed_experts_if_allowed, + select_experts, +) if TYPE_CHECKING: from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo @@ -104,10 +107,6 @@ def fused_topk_npu( if expert_location_dispatch_info is not None: topk_ids = topk_ids_logical_to_physical(topk_ids, expert_location_dispatch_info) get_global_expert_distribution_recorder().on_select_experts(topk_ids=topk_ids) - if (cap := get_global_experts_capturer()) is not None: - cap.capture( - layer_id=layer_id, - topk_indices=topk_ids, - ) + capture_routed_experts_if_allowed(topk_config, layer_id, topk_ids) return StandardTopKOutput(topk_weights, topk_ids, router_logits) diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 6acd20ed0..db10d3ce4 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -1103,6 +1103,7 @@ class FusedMoE(torch.nn.Module): topk_output.topk_config.correction_bias, topk_output.topk_config.renormalize, self.layer_id, + topk_output.topk_config.allow_routed_experts_capture, ) else: # Make sure there is torch lib op registration for the whole moe layer @@ -1353,6 +1354,7 @@ def fused_moe_bypassed_piecewise_cuda_graph_impl( correction_bias: Optional[torch.Tensor], renormalize: bool, layer_id: int, + allow_routed_experts_capture: bool, ) -> torch.Tensor: topk_output = BypassedTopKOutput( hidden_states=hidden_states, @@ -1363,6 +1365,7 @@ def fused_moe_bypassed_piecewise_cuda_graph_impl( num_expert_group=num_expert_group, correction_bias=correction_bias, renormalize=renormalize, + allow_routed_experts_capture=allow_routed_experts_capture, ), ) forward_context = get_tc_piecewise_forward_context() diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index f8d9df3a2..40fe623f1 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -222,6 +222,9 @@ class TopKConfig: fused_shared_experts_scaling_factor: Optional[float] = None output_format: Optional[TopKOutputFormat] = None scoring_func: str = "softmax" + # Draft-side MoE blocks set this False so they never write the target's + # process-global routed-experts capture buffer. + allow_routed_experts_capture: bool = True # -------------------------------- TopKOutput --------------------------------------- @@ -386,6 +389,7 @@ class TopK(MultiPlatformOp): output_format: Optional[TopKOutputFormat] = None, fused_shared_experts_scaling_factor: Optional[float] = None, is_fp4_experts: bool = False, + allow_routed_experts_capture: bool = True, ): # NOTE: scoring_func is not used for now, but we keep it for future use # see https://github.com/sgl-project/sglang/pull/4505 for more details @@ -426,6 +430,7 @@ class TopK(MultiPlatformOp): fused_shared_experts_scaling_factor=fused_shared_experts_scaling_factor, output_format=output_format, scoring_func=scoring_func, + allow_routed_experts_capture=allow_routed_experts_capture, ) def _apply_deepep_waterfill( @@ -1597,6 +1602,25 @@ def _remap_topk_for_deepep( return topk_ids, topk_weights +def capture_routed_experts_if_allowed( + topk_config: TopKConfig, + layer_id: Optional[int], + topk_ids: torch.Tensor, +) -> None: + """Single capture site for every backend, gated by the per-config opt-out. + + Routing all backends through here keeps the draft-side opt-out from being + bypassed by an inlined capturer call. + """ + if not topk_config.allow_routed_experts_capture: + return + if (cap := get_global_experts_capturer()) is not None: + cap.capture( + layer_id=layer_id, + topk_indices=topk_ids, + ) + + def _post_process_topk_ids( topk_ids: torch.Tensor, topk_weights: torch.Tensor, @@ -1610,11 +1634,7 @@ def _post_process_topk_ids( fused_shared_experts_scaling_factor = ( topk_config.fused_shared_experts_scaling_factor ) - if (cap := get_global_experts_capturer()) is not None: - cap.capture( - layer_id=layer_id, - topk_indices=topk_ids, - ) + capture_routed_experts_if_allowed(topk_config, layer_id, topk_ids) recorder_topk_ids = None if _is_cuda: # LP path: solve LP outside torch.compile (the solver contains an diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 3c77948da..afcfafd36 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -192,6 +192,7 @@ from sglang.srt.state_capturer.indexer_topk import ( ) from sglang.srt.state_capturer.routed_experts import ( RoutedExpertsCapturer, + disable_routed_experts_capture_for_draft, get_global_experts_capturer, set_global_experts_capturer, ) @@ -699,6 +700,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.load_model() self._prepare_moe_topk() + # Must run before backend/graph init so no draft graph records a + # routed-experts capture-write kernel. + if self.is_draft_worker: + disable_routed_experts_capture_for_draft(self.model) + # Load the expert backup client self.expert_backup_client = ( ExpertBackupClient(self.server_args, self) @@ -960,6 +966,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.model_config.full_attention_layer_ids = full_attention_layer_ids def init_routed_experts_capturer(self): + if self.is_draft_worker: + # Capture is target-only. The draft worker runs in the same process + # as its target and inits after it, so installing a capturer here + # would overwrite the target's process-global one. + return + if not self.server_args.disable_shared_experts_fusion and hasattr( self.model, "num_fused_shared_experts" ): @@ -2960,7 +2972,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): output.expert_distribution_metrics = recorder_outputs.get("metrics") no_copy_to_cpu = not self.server_args.disable_overlap_schedule - if (experts_capturer := get_global_experts_capturer()) is not None: + if ( + not self.is_draft_worker + and (experts_capturer := get_global_experts_capturer()) is not None + ): output.routed_experts_output = experts_capturer.on_forward_end( forward_batch=forward_batch, can_run_graph=output.can_run_graph, diff --git a/python/sglang/srt/state_capturer/routed_experts.py b/python/sglang/srt/state_capturer/routed_experts.py index 44f339e00..982cd4a49 100644 --- a/python/sglang/srt/state_capturer/routed_experts.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -1,4 +1,4 @@ -from typing import Optional +from typing import Any, Optional import numpy as np import pybase64 @@ -146,3 +146,19 @@ def extract_routed_experts_from_meta_info(data): pybase64.b64decode(routed_experts_base64.encode("utf-8")), dtype=np.int32 ) return routed_experts + + +def disable_routed_experts_capture_for_draft(model: Any) -> None: + """Opt every draft MoE ``TopK`` out of routed-experts (R3) capture. + + Capture is target-only; a draft ``TopK`` must never write the target's + process-global buffer. ``HashTopK`` has no ``topk_config`` and never + captures, so it is left untouched. + """ + # Lazy import: ``layers.moe.topk`` imports ``get_global_experts_capturer`` + # from this module, so a top-level import here would be circular. + from sglang.srt.layers.moe.topk import TopK + + for module in model.modules(): + if isinstance(module, TopK): + module.topk_config.allow_routed_experts_capture = False