[fix] Skip routed expert capture for draft model under spec v2 (#26980)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
7e63feee6f
commit
e4bf0043fe
@@ -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_distribution import get_global_expert_distribution_recorder
|
||||||
from sglang.srt.eplb.expert_location_dispatch import topk_ids_logical_to_physical
|
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.layers.moe.topk import (
|
||||||
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
|
StandardTopKOutput,
|
||||||
|
capture_routed_experts_if_allowed,
|
||||||
|
select_experts,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
|
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:
|
if expert_location_dispatch_info is not None:
|
||||||
topk_ids = topk_ids_logical_to_physical(topk_ids, expert_location_dispatch_info)
|
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)
|
get_global_expert_distribution_recorder().on_select_experts(topk_ids=topk_ids)
|
||||||
if (cap := get_global_experts_capturer()) is not None:
|
capture_routed_experts_if_allowed(topk_config, layer_id, topk_ids)
|
||||||
cap.capture(
|
|
||||||
layer_id=layer_id,
|
|
||||||
topk_indices=topk_ids,
|
|
||||||
)
|
|
||||||
|
|
||||||
return StandardTopKOutput(topk_weights, topk_ids, router_logits)
|
return StandardTopKOutput(topk_weights, topk_ids, router_logits)
|
||||||
|
|||||||
@@ -1103,6 +1103,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
topk_output.topk_config.correction_bias,
|
topk_output.topk_config.correction_bias,
|
||||||
topk_output.topk_config.renormalize,
|
topk_output.topk_config.renormalize,
|
||||||
self.layer_id,
|
self.layer_id,
|
||||||
|
topk_output.topk_config.allow_routed_experts_capture,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Make sure there is torch lib op registration for the whole moe layer
|
# 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],
|
correction_bias: Optional[torch.Tensor],
|
||||||
renormalize: bool,
|
renormalize: bool,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
|
allow_routed_experts_capture: bool,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
topk_output = BypassedTopKOutput(
|
topk_output = BypassedTopKOutput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
@@ -1363,6 +1365,7 @@ def fused_moe_bypassed_piecewise_cuda_graph_impl(
|
|||||||
num_expert_group=num_expert_group,
|
num_expert_group=num_expert_group,
|
||||||
correction_bias=correction_bias,
|
correction_bias=correction_bias,
|
||||||
renormalize=renormalize,
|
renormalize=renormalize,
|
||||||
|
allow_routed_experts_capture=allow_routed_experts_capture,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
forward_context = get_tc_piecewise_forward_context()
|
forward_context = get_tc_piecewise_forward_context()
|
||||||
|
|||||||
@@ -222,6 +222,9 @@ class TopKConfig:
|
|||||||
fused_shared_experts_scaling_factor: Optional[float] = None
|
fused_shared_experts_scaling_factor: Optional[float] = None
|
||||||
output_format: Optional[TopKOutputFormat] = None
|
output_format: Optional[TopKOutputFormat] = None
|
||||||
scoring_func: str = "softmax"
|
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 ---------------------------------------
|
# -------------------------------- TopKOutput ---------------------------------------
|
||||||
@@ -386,6 +389,7 @@ class TopK(MultiPlatformOp):
|
|||||||
output_format: Optional[TopKOutputFormat] = None,
|
output_format: Optional[TopKOutputFormat] = None,
|
||||||
fused_shared_experts_scaling_factor: Optional[float] = None,
|
fused_shared_experts_scaling_factor: Optional[float] = None,
|
||||||
is_fp4_experts: bool = False,
|
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
|
# 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
|
# 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,
|
fused_shared_experts_scaling_factor=fused_shared_experts_scaling_factor,
|
||||||
output_format=output_format,
|
output_format=output_format,
|
||||||
scoring_func=scoring_func,
|
scoring_func=scoring_func,
|
||||||
|
allow_routed_experts_capture=allow_routed_experts_capture,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _apply_deepep_waterfill(
|
def _apply_deepep_waterfill(
|
||||||
@@ -1597,6 +1602,25 @@ def _remap_topk_for_deepep(
|
|||||||
return topk_ids, topk_weights
|
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(
|
def _post_process_topk_ids(
|
||||||
topk_ids: torch.Tensor,
|
topk_ids: torch.Tensor,
|
||||||
topk_weights: torch.Tensor,
|
topk_weights: torch.Tensor,
|
||||||
@@ -1610,11 +1634,7 @@ def _post_process_topk_ids(
|
|||||||
fused_shared_experts_scaling_factor = (
|
fused_shared_experts_scaling_factor = (
|
||||||
topk_config.fused_shared_experts_scaling_factor
|
topk_config.fused_shared_experts_scaling_factor
|
||||||
)
|
)
|
||||||
if (cap := get_global_experts_capturer()) is not None:
|
capture_routed_experts_if_allowed(topk_config, layer_id, topk_ids)
|
||||||
cap.capture(
|
|
||||||
layer_id=layer_id,
|
|
||||||
topk_indices=topk_ids,
|
|
||||||
)
|
|
||||||
recorder_topk_ids = None
|
recorder_topk_ids = None
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
# LP path: solve LP outside torch.compile (the solver contains an
|
# LP path: solve LP outside torch.compile (the solver contains an
|
||||||
|
|||||||
@@ -192,6 +192,7 @@ from sglang.srt.state_capturer.indexer_topk import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.state_capturer.routed_experts import (
|
from sglang.srt.state_capturer.routed_experts import (
|
||||||
RoutedExpertsCapturer,
|
RoutedExpertsCapturer,
|
||||||
|
disable_routed_experts_capture_for_draft,
|
||||||
get_global_experts_capturer,
|
get_global_experts_capturer,
|
||||||
set_global_experts_capturer,
|
set_global_experts_capturer,
|
||||||
)
|
)
|
||||||
@@ -699,6 +700,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.load_model()
|
self.load_model()
|
||||||
self._prepare_moe_topk()
|
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
|
# Load the expert backup client
|
||||||
self.expert_backup_client = (
|
self.expert_backup_client = (
|
||||||
ExpertBackupClient(self.server_args, self)
|
ExpertBackupClient(self.server_args, self)
|
||||||
@@ -960,6 +966,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.model_config.full_attention_layer_ids = full_attention_layer_ids
|
self.model_config.full_attention_layer_ids = full_attention_layer_ids
|
||||||
|
|
||||||
def init_routed_experts_capturer(self):
|
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(
|
if not self.server_args.disable_shared_experts_fusion and hasattr(
|
||||||
self.model, "num_fused_shared_experts"
|
self.model, "num_fused_shared_experts"
|
||||||
):
|
):
|
||||||
@@ -2960,7 +2972,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
output.expert_distribution_metrics = recorder_outputs.get("metrics")
|
output.expert_distribution_metrics = recorder_outputs.get("metrics")
|
||||||
|
|
||||||
no_copy_to_cpu = not self.server_args.disable_overlap_schedule
|
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(
|
output.routed_experts_output = experts_capturer.on_forward_end(
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
can_run_graph=output.can_run_graph,
|
can_run_graph=output.can_run_graph,
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pybase64
|
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
|
pybase64.b64decode(routed_experts_base64.encode("utf-8")), dtype=np.int32
|
||||||
)
|
)
|
||||||
return routed_experts
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user