[perf] support return_routed_experts with overlap scheduling (#22911)

Co-authored-by: Yuzhen Zhou <82826991+zyzshishui@users.noreply.github.com>
This commit is contained in:
Qiaolin Yu
2026-04-21 14:42:49 -07:00
committed by GitHub
co-authored by Yuzhen Zhou
parent 9f37c1a9b0
commit c560326884
7 changed files with 93 additions and 24 deletions
@@ -1,3 +1,4 @@
import dataclasses
import logging
from abc import ABC
from typing import Optional
@@ -26,6 +27,25 @@ def get_tensor_size_bytes(t: torch.Tensor):
return np.prod(t.shape) * t.dtype.itemsize
@dataclasses.dataclass
class RoutedExpertsOutput:
"""Holds GPU tensors captured during forward for overlap scheduling.
Call copy_to_cpu() inside forward stream (before copy_done.record()),
then finalize() after copy_done.synchronize().
"""
out_cache_loc: torch.Tensor
routed_experts: torch.Tensor
host_cache: "_RoutedExpertsHostCache"
def copy_to_cpu(self):
self.out_cache_loc = self.out_cache_loc.to("cpu", non_blocking=True)
self.routed_experts = self.routed_experts.to("cpu", non_blocking=True)
def finalize(self):
self.host_cache.buffer[self.out_cache_loc] = self.routed_experts
class _RoutedExpertsDeviceCache:
def __init__(
self,
@@ -142,7 +162,9 @@ class RoutedExpertsCapturer(ABC):
):
raise NotImplementedError
def on_forward_end(self, forward_batch, can_run_graph, cuda_graph_batch):
def on_forward_end(
self, forward_batch, can_run_graph, cuda_graph_batch, no_copy_to_cpu=False
) -> Optional[RoutedExpertsOutput]:
raise NotImplementedError
def get_host_cache(self):
@@ -181,30 +203,46 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer):
device=device,
)
def _get_local_range(self, forward_batch, can_run_graph, cuda_graph_batch):
if is_dp_attention_enabled():
local_start_pos, local_num_tokens = get_dp_local_info(forward_batch)
if can_run_graph:
local_start_pos = get_attention_dp_rank() * cuda_graph_batch
return local_start_pos, local_start_pos + local_num_tokens
else:
return 0, forward_batch.out_cache_loc.shape[0]
def _sync_fwd_experts_buffer_DtoH(
self,
forward_batch: ForwardBatch,
can_run_graph: bool,
cuda_graph_batch: int,
):
if is_dp_attention_enabled():
local_start_pos, local_num_tokens = get_dp_local_info(forward_batch)
# handle with cuda graph padding
if can_run_graph:
local_start_pos = get_attention_dp_rank() * cuda_graph_batch
local_end_pos = local_start_pos + local_num_tokens
else:
local_end_pos = local_start_pos + local_num_tokens
else:
local_start_pos = 0
local_end_pos = forward_batch.out_cache_loc.shape[0]
# FIXME: sync explicitly here, overlap scheduler breaks here.
local_start_pos, local_end_pos = self._get_local_range(
forward_batch, can_run_graph, cuda_graph_batch
)
out_cache_loc_cpu = forward_batch.out_cache_loc.cpu()
self.host_cache.buffer[out_cache_loc_cpu] = self.device_cache.buffer[
local_start_pos:local_end_pos, :, : self.num_experts_per_tok
].cpu()
def _prepare_routed_experts_output(
self,
forward_batch: ForwardBatch,
can_run_graph: bool,
cuda_graph_batch: int,
) -> RoutedExpertsOutput:
local_start_pos, local_end_pos = self._get_local_range(
forward_batch, can_run_graph, cuda_graph_batch
)
return RoutedExpertsOutput(
out_cache_loc=forward_batch.out_cache_loc,
routed_experts=self.device_cache.buffer[
local_start_pos:local_end_pos, :, : self.num_experts_per_tok
],
host_cache=self.host_cache,
)
def capture(self, layer_id: int, topk_ids: torch.Tensor):
self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids)
@@ -219,12 +257,22 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer):
)
return self.get_host_cache().buffer[cache_pool_idx]
def on_forward_end(self, forward_batch, can_run_graph, cuda_graph_batch):
self._sync_fwd_experts_buffer_DtoH(
forward_batch=forward_batch,
can_run_graph=can_run_graph,
cuda_graph_batch=cuda_graph_batch,
)
def on_forward_end(
self, forward_batch, can_run_graph, cuda_graph_batch, no_copy_to_cpu=False
) -> Optional[RoutedExpertsOutput]:
if no_copy_to_cpu:
return self._prepare_routed_experts_output(
forward_batch=forward_batch,
can_run_graph=can_run_graph,
cuda_graph_batch=cuda_graph_batch,
)
else:
self._sync_fwd_experts_buffer_DtoH(
forward_batch=forward_batch,
can_run_graph=can_run_graph,
cuda_graph_batch=cuda_graph_batch,
)
return None
def get_host_cache(self):
return self.host_cache
@@ -256,8 +304,10 @@ class _RoutedExpertsCapturerNoop(RoutedExpertsCapturer):
):
pass
def on_forward_end(self, forward_batch, can_run_graph, cuda_graph_batch):
pass
def on_forward_end(
self, forward_batch, can_run_graph, cuda_graph_batch, no_copy_to_cpu=False
) -> Optional[RoutedExpertsOutput]:
return None
def get_host_cache(self):
pass
@@ -133,6 +133,9 @@ class SchedulerOutputProcessorMixin:
if self.is_generation:
if result.copy_done is not None:
result.copy_done.synchronize()
if result.routed_experts_output is not None:
result.routed_experts_output.finalize()
result.routed_experts_output = None
(
logits_output,
@@ -391,6 +394,9 @@ class SchedulerOutputProcessorMixin:
):
if result.copy_done is not None:
result.copy_done.synchronize()
if result.routed_experts_output is not None:
result.routed_experts_output.finalize()
result.routed_experts_output = None
logits_output, next_token_ids, can_run_cuda_graph = (
result.logits_output,
+1
View File
@@ -475,6 +475,7 @@ class TpModelWorker(BaseTpWorker):
logits_output=logits_output,
can_run_cuda_graph=can_run_cuda_graph,
expert_distribution_metrics=out.expert_distribution_metrics,
routed_experts_output=out.routed_experts_output,
)
if is_verify:
+7
View File
@@ -8,6 +8,7 @@ import torch
from sglang.srt.eplb.expert_distribution import ExpertDistributionMetrics
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.routed_experts_capturer import RoutedExpertsOutput
from sglang.srt.managers.overlap_utils import FutureIndices
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
@@ -46,6 +47,9 @@ class GenerationBatchResult:
# relay path: forward stream -> next step forward
next_draft_input: Optional[EagleDraftInput] = None
# Routed experts: pending async D2H for overlap scheduling
routed_experts_output: Optional[RoutedExpertsOutput] = None
# metrics
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None
@@ -87,6 +91,9 @@ class GenerationBatchResult:
if self.accept_lens is not None:
self.accept_lens = self.accept_lens.to("cpu", non_blocking=True)
if self.routed_experts_output is not None:
self.routed_experts_output.copy_to_cpu()
if (x := self.expert_distribution_metrics) is not None:
x.copy_to_cpu()
@@ -110,6 +110,7 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.routed_experts_capturer import (
RoutedExpertsCapturer,
RoutedExpertsOutput,
get_global_experts_capturer,
set_global_experts_capturer,
)
@@ -287,6 +288,7 @@ class ModelRunnerOutput:
logits_output: Union[LogitsProcessorOutput, PPProxyTensors]
can_run_graph: bool
expert_distribution_metrics: Optional[ExpertDistributionMetrics] = None
routed_experts_output: Optional[RoutedExpertsOutput] = None
class ModelRunner(ModelRunnerKVCacheMixin):
@@ -2934,11 +2936,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
)
output.expert_distribution_metrics = recorder_outputs.get("metrics")
# Copy cached routing experts' buffers back to CPU cache
get_global_experts_capturer().on_forward_end(
no_copy_to_cpu = not self.server_args.disable_overlap_schedule
output.routed_experts_output = get_global_experts_capturer().on_forward_end(
forward_batch=forward_batch,
can_run_graph=output.can_run_graph,
cuda_graph_batch=getattr(self.graph_runner, "bs", None),
no_copy_to_cpu=no_copy_to_cpu,
)
if self.eplb_manager is not None:
@@ -872,6 +872,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
can_run_cuda_graph=can_run_cuda_graph,
next_draft_input=next_draft_input,
accept_lens=accept_length,
routed_experts_output=forward_batch_output.routed_experts_output,
)
def _mamba_verify_update(
@@ -781,6 +781,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
can_run_cuda_graph=can_run_cuda_graph,
next_draft_input=next_draft_input,
accept_lens=accept_length,
routed_experts_output=forward_batch_output.routed_experts_output,
)
def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput):