[Feature] Enable return routed experts (#12162)
Co-authored-by: yizhang2077 <1109276519@qq.com> Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
This commit is contained in:
@@ -97,6 +97,11 @@ from sglang.srt.layers.dp_attention import (
|
||||
set_is_extend_in_batch,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.moe.routed_experts_capturer import (
|
||||
RoutedExpertsCapturer,
|
||||
get_global_experts_capturer,
|
||||
set_global_experts_capturer,
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import get_moe_a2a_backend
|
||||
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
||||
@@ -557,6 +562,21 @@ class ModelRunner:
|
||||
server_args.max_running_requests,
|
||||
server_args.max_total_tokens,
|
||||
)
|
||||
|
||||
# Init max running requests
|
||||
self.max_running_requests = min(
|
||||
(
|
||||
self.max_total_num_tokens // 2
|
||||
if server_args.max_running_requests is None
|
||||
else server_args.max_running_requests
|
||||
// (server_args.dp_size if server_args.enable_dp_attention else 1)
|
||||
),
|
||||
self.req_to_token_pool.size,
|
||||
)
|
||||
|
||||
# Init routed experts capturer
|
||||
self.init_routed_experts_capturer()
|
||||
|
||||
if self.device == "cuda":
|
||||
self.init_cublas()
|
||||
self.init_attention_backend()
|
||||
@@ -600,6 +620,40 @@ class ModelRunner:
|
||||
# Initialize piecewise CUDA graph
|
||||
self.init_piecewise_cuda_graphs()
|
||||
|
||||
def init_routed_experts_capturer(self):
|
||||
# TODO: the redundant logic with TpModelWorker
|
||||
max_running_requests = min(
|
||||
(
|
||||
self.max_total_num_tokens // 2
|
||||
if self.server_args.max_running_requests is None
|
||||
else self.server_args.max_running_requests
|
||||
// (
|
||||
self.server_args.dp_size
|
||||
if self.server_args.enable_dp_attention
|
||||
else 1
|
||||
)
|
||||
),
|
||||
self.req_to_token_pool.size,
|
||||
)
|
||||
|
||||
if not self.server_args.disable_shared_experts_fusion and hasattr(
|
||||
self.model, "num_fused_shared_experts"
|
||||
):
|
||||
num_fused_shared_experts = self.model.num_fused_shared_experts
|
||||
else:
|
||||
num_fused_shared_experts = 0
|
||||
|
||||
set_global_experts_capturer(
|
||||
RoutedExpertsCapturer.create(
|
||||
enable=get_global_server_args().enable_return_routed_experts,
|
||||
model_config=self.model_config,
|
||||
num_fused_shared_experts=num_fused_shared_experts,
|
||||
num_tokens=self.max_total_num_tokens + self.page_size,
|
||||
max_running_requests=max_running_requests,
|
||||
device=self.device,
|
||||
)
|
||||
)
|
||||
|
||||
def remote_instance_init_transfer_engine(self):
|
||||
try:
|
||||
from mooncake.engine import TransferEngine
|
||||
@@ -2840,6 +2894,13 @@ class ModelRunner:
|
||||
)
|
||||
output.expert_distribution_metrics = recorder_outputs.get("metrics")
|
||||
|
||||
# Copy cached routing experts' buffers back to CPU cache
|
||||
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),
|
||||
)
|
||||
|
||||
if self.eplb_manager is not None:
|
||||
self.eplb_manager.on_forward_pass_end()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user