Register allgather/reducescatter buffers with symm memory (#12572)

This commit is contained in:
Nicolas Castet
2025-11-04 17:11:36 -08:00
committed by GitHub
parent 1357ab025a
commit 2340798353
19 changed files with 250 additions and 114 deletions
+21 -9
View File
@@ -32,12 +32,17 @@ import torch
import torch.nn.functional as F
from sglang.srt.custom_op import CustomOp
from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.eplb import expert_location_dispatch
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location_dispatch import (
ExpertLocationDispatchInfo,
topk_ids_logical_to_physical,
)
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe import get_moe_runner_backend
from sglang.srt.utils import (
cpu_has_amx_support,
@@ -279,13 +284,17 @@ class TopK(CustomOp):
)
else:
self.topk_config.torch_native = False
return select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
topk_config=self.topk_config,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
topk_output = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
topk_config=self.topk_config,
num_token_non_padded=num_token_non_padded,
expert_location_dispatch_info=expert_location_dispatch_info,
)
return topk_output
def forward_cpu(
self,
@@ -386,8 +395,11 @@ class TopK(CustomOp):
def empty_topk_output(self, device: torch.device) -> TopKOutput:
topk = self.topk_config.top_k - self.topk_config.num_fused_shared_experts
topk_weights = torch.empty((0, topk), dtype=torch.float32, device=device)
topk_ids = torch.full((0, topk), -1, dtype=torch.int32, device=device)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
topk_weights = torch.empty((0, topk), dtype=torch.float32, device=device)
topk_ids = torch.full((0, topk), -1, dtype=torch.int32, device=device)
# FIXME: router_logits should be of size (0, num_experts)
router_logits = torch.empty((0, topk), dtype=torch.float32, device=device)
return StandardTopKOutput(topk_weights, topk_ids, router_logits)