Register allgather/reducescatter buffers with symm memory (#12572)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user