[Feature] Add DWDP (Distributed Weight Data Parallelism) for MoE prefill (#29778)

This commit is contained in:
Yuhao Yang
2026-07-20 23:59:54 -07:00
committed by GitHub
parent 01f558d905
commit 37a830b667
20 changed files with 2147 additions and 24 deletions
@@ -50,6 +50,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
is_in_tc_piecewise_cuda_graph,
)
from sglang.srt.platforms.device_mixin import _DEVICE_TO_DISTRIBUTED_BACKEND
from sglang.srt.runtime_context import (
get_global_dwdp_manager,
set_global_dwdp_manager,
)
from sglang.srt.utils import (
get_current_device_stream_fast,
get_int_env_var,
@@ -2615,6 +2619,11 @@ def get_moe_tensor_parallel_rank():
def destroy_model_parallel():
"""Set the groups to none and destroy them."""
dwdp_mgr = get_global_dwdp_manager()
if dwdp_mgr is not None:
dwdp_mgr.cleanup()
set_global_dwdp_manager(None)
global _TP
if _TP:
_TP.destroy()
+5
View File
@@ -385,6 +385,7 @@ class LayerScatterModes:
# Token dispatch/combine will be handled outside of LayerCommunicator for these modes.
not get_moe_a2a_backend().is_none()
or should_use_flashinfer_cutlass_moe_fp4_allgather()
or enable_dwdp()
):
return ScatterMode.SCATTERED
# DSA CP and MLA CP both don't support MOE_FULL yet; fall back to FULL.
@@ -436,6 +437,10 @@ def enable_moe_dense_fully_dp():
return get_server_args().moe_dense_tp_size == 1
def enable_dwdp():
return get_server_args().dwdp_size > 1
class LayerCommunicator:
def __init__(
self,
@@ -0,0 +1,13 @@
"""DWDP (Distributed Weight Data Parallelism): MoE prefill with tokens kept on-rank and peer expert weights prefetched via NVLink into a composite VMM address space."""
from sglang.srt.layers.moe.dwdp.dwdp_manager import DwdpManager
from sglang.srt.runtime_context import (
get_global_dwdp_manager,
set_global_dwdp_manager,
)
__all__ = [
"DwdpManager",
"get_global_dwdp_manager",
"set_global_dwdp_manager",
]
@@ -0,0 +1,211 @@
"""Global singleton orchestrating the DWDP lifecycle from setup(model) to cleanup()."""
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import torch
import torch.distributed as dist
from torch import nn
from sglang.srt.layers.moe.dwdp.layout import (
DwdpExpertLayout,
build_layer_weight_specs,
lookup_owner,
)
from sglang.srt.layers.moe.dwdp.transport import DWDPTransport
from sglang.srt.layers.moe.dwdp.weight_buffer import WeightBuffer
from sglang.srt.layers.moe.dwdp.weight_manager import DWDPWeightManager
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
_EXPERT_WEIGHT_NAMES = (
"w13_weight",
"w2_weight",
)
class DwdpManager:
def __init__(self, server_args: ServerArgs):
self.dwdp_size = server_args.dwdp_size
self.dwdp_rank = get_parallel().tp_rank
self.device_id = torch.cuda.current_device()
self.layout: Optional[DwdpExpertLayout] = None
self._weight_manager: Optional[DWDPWeightManager] = None
self._moe_layer_indices: List[int] = []
def setup(self, model: nn.Module) -> None:
if self._weight_manager is not None:
return
moe_layers = self._collect_moe_layers(model)
if not moe_layers:
raise RuntimeError(
f"DWDP is enabled but no FusedMoE layers were found in "
f"{type(model).__name__}"
)
self._moe_layer_indices = [li for li, _ in moe_layers]
expert_counts = {e.num_global_routed_experts for _, e in moe_layers}
if len(expert_counts) != 1:
raise RuntimeError(
f"DWDP requires a uniform routed expert count across MoE layers, "
f"got {sorted(expert_counts)}"
)
num_routed = expert_counts.pop()
if num_routed % self.dwdp_size != 0:
raise ValueError(
f"DWDP requires num_routed_experts ({num_routed}) to be divisible "
f"by dwdp_size ({self.dwdp_size})"
)
self.layout = DwdpExpertLayout(
num_routed_experts=num_routed,
dwdp_size=self.dwdp_size,
dwdp_rank=self.dwdp_rank,
)
logger.info(
f"DWDP layout: {self.layout.num_routed_experts} experts, "
f"local [{self.layout.local_expert_start}, {self.layout.local_expert_end}), "
f"prefetch_per_peer={self.layout.num_prefetch_experts}"
)
local_params = {}
for li, experts in moe_layers:
local_params[(li, "w13_weight")] = experts.w13_weight.data
local_params[(li, "w2_weight")] = experts.w2_weight.data
layer_weight_specs = build_layer_weight_specs(
local_params, self.layout.num_routed_experts
)
group = get_parallel().tp_group
transport = DWDPTransport.create(
layer_weight_specs=layer_weight_specs,
local_params=local_params,
group=group,
layout=self.layout,
device_id=self.device_id,
)
weight_buffer = WeightBuffer.create(
layer_weight_specs=layer_weight_specs,
handles=transport.handle_set,
local_start=self.layout.local_expert_start,
local_end=self.layout.local_expert_end,
dwdp_size=self.dwdp_size,
device_id=self.device_id,
)
self._fill_edge_bytes(weight_buffer, transport.peer_views)
self._weight_manager = DWDPWeightManager(
weight_buffer=weight_buffer,
peer_views=transport.peer_views,
peer_ranges=self.layout.peer_ranges,
moe_layer_indices=self._moe_layer_indices,
weight_names=list(_EXPERT_WEIGHT_NAMES),
dwdp_rank=self.dwdp_rank,
dwdp_size=self.dwdp_size,
transport=transport,
)
for li, experts in moe_layers:
experts.bind_full_expert_weights(
{
name: weight_buffer.get_full_tensor(li, name)
for name in weight_buffer.weight_names(li)
}
)
self._allgather_small_params(moe_layers, group)
logger.info("DWDP setup complete.")
def prefetch_first_layers(self) -> None:
if self._weight_manager is not None:
self._weight_manager.prefetch_first_layers()
def wait_prefetch(self, layer_idx: int) -> None:
if self._weight_manager is not None:
self._weight_manager.wait_prefetch(layer_idx)
def record_compute_and_prefetch_next(self, layer_idx: int) -> None:
if self._weight_manager is not None:
self._weight_manager.record_compute_and_prefetch_next(layer_idx)
def cleanup(self) -> None:
if self._weight_manager is not None:
self._weight_manager.release()
self._weight_manager = None
@staticmethod
def _collect_moe_layers(model: nn.Module) -> List[Tuple[int, FusedMoE]]:
decoder = model.model if hasattr(model, "model") else model
moe_layers = []
for layer_idx, layer in enumerate(decoder.layers):
experts = next(
(m for m in layer.modules() if isinstance(m, FusedMoE)), None
)
if experts is not None:
moe_layers.append((layer_idx, experts))
return moe_layers
def _fill_edge_bytes(
self,
weight_buffer: WeightBuffer,
peer_views: Dict[Tuple[int, int, str], torch.Tensor],
) -> None:
local_start = self.layout.local_expert_start
local_end = self.layout.local_expert_end
peer_ranges = self.layout.peer_ranges
for li in weight_buffer.layer_indices:
for name in weight_buffer.weight_names(li):
edge = weight_buffer.get_edge_info(li, name)
if edge.leading_edge == 0 and edge.trailing_edge == 0:
continue
full_tensor = weight_buffer.get_full_tensor(li, name)
if edge.leading_edge > 0 and local_start > 0:
prev = local_start - 1
peer = lookup_owner(prev, peer_ranges)
ps, _ = peer_ranges[peer]
key = (peer, li, name)
if key in peer_views:
full_tensor[prev].copy_(peer_views[key][prev - ps])
if edge.trailing_edge > 0 and local_end < full_tensor.shape[0]:
nxt = local_end
peer = lookup_owner(nxt, peer_ranges)
ps, _ = peer_ranges[peer]
key = (peer, li, name)
if key in peer_views:
full_tensor[nxt].copy_(peer_views[key][nxt - ps])
torch.cuda.synchronize(weight_buffer.device_id)
def _allgather_small_params(
self, moe_layers: List[Tuple[int, FusedMoE]], group
) -> None:
local_experts = self.layout.num_experts_per_worker
num_total = self.layout.num_routed_experts
for li, experts in moe_layers:
for pname, data in experts.named_per_expert_tensors(local_experts):
shards = [torch.empty_like(data) for _ in range(self.dwdp_size)]
dist.all_gather(shards, data, group=group.device_group)
full = torch.cat(shards, dim=0)[:num_total].contiguous()
experts.replace_expert_tensor(pname, full)
logger.debug(
f"Layer {li}: allgathered {pname} "
f"({local_experts} -> {full.shape[0]}) "
f"shape={tuple(full.shape)} dtype={full.dtype} "
f"size={full.numel() * full.element_size() / 1e6:.1f}MB"
)
+287
View File
@@ -0,0 +1,287 @@
# Adapted from NVIDIA TensorRT-LLM (https://github.com/NVIDIA/TensorRT-LLM)
"""Expert ownership and page-aligned memory layout computation for DWDP."""
from __future__ import annotations
import math
from typing import Dict, List, Optional, Tuple
import torch
from sglang.srt.layers.moe.dwdp.vmm import align_down, align_up
# one (start, end_capped) expert range per peer DWDP rank
PeerRanges = List[Tuple[int, int]]
LayerWeightSpecs = Dict[int, Dict[str, "WeightSpec"]]
class DwdpExpertLayout:
def __init__(
self,
num_routed_experts: int,
dwdp_size: int,
dwdp_rank: int,
):
self.num_routed_experts = num_routed_experts
self.dwdp_size = dwdp_size
self.dwdp_rank = dwdp_rank
num_experts_per_worker = num_routed_experts // dwdp_size
self.num_experts_per_worker = num_experts_per_worker
self.num_prefetch_experts = math.ceil(
(num_routed_experts - num_experts_per_worker) / (dwdp_size - 1)
)
self.local_expert_start = min(
self.num_prefetch_experts * dwdp_rank,
num_routed_experts - num_experts_per_worker,
)
self.local_expert_end = self.local_expert_start + num_experts_per_worker
self.peer_ranges = compute_peer_ranges(
dwdp_size=dwdp_size,
num_experts_per_worker=num_experts_per_worker,
num_prefetch_experts=self.num_prefetch_experts,
num_experts_total=num_routed_experts,
)
class WeightSpec:
def __init__(
self,
num_experts: int,
chunk_shape: Tuple[int, ...],
full_shape: Tuple[int, ...],
dtype: torch.dtype,
):
self.num_experts = num_experts
self.chunk_shape = chunk_shape
self.full_shape = full_shape
self.dtype = dtype
@property
def expert_bytes(self) -> int:
n = 1
for d in self.full_shape[1:]:
n *= d
return n * torch.tensor([], dtype=self.dtype).element_size()
@property
def chunk_bytes(self) -> int:
n = 1
for d in self.chunk_shape:
n *= d
return n * torch.tensor([], dtype=self.dtype).element_size()
@property
def local_experts(self) -> int:
return self.chunk_shape[0]
class EdgeInfo:
def __init__(
self,
data_offset: int,
leading_edge: int,
trailing_edge: int,
page_start: int,
page_end: int,
expert_bytes: int,
):
self.data_offset = data_offset
self.leading_edge = leading_edge
self.trailing_edge = trailing_edge
self.page_start = page_start
self.page_end = page_end
self.expert_bytes = expert_bytes
class PageAlignedLayout:
# composite VA: [pre_region (pool pages) | mnnvl_region (fabric handle) | post_region (pool pages)]
def __init__(
self,
expert_bytes: int,
num_experts: int,
local_start: int,
local_end: int,
granularity: int,
pool_granularity: int,
page_start: int,
page_end: int,
pre_size: int,
mnnvl_size: int,
post_size: int,
pre_padding: int,
post_padding: int,
data_offset: int,
leading_edge: int,
trailing_edge: int,
total_size: int,
handle_phys_size: int,
):
self.expert_bytes = expert_bytes
self.num_experts = num_experts
self.local_start = local_start
self.local_end = local_end
self.granularity = granularity
self.pool_granularity = pool_granularity
self.page_start = page_start
self.page_end = page_end
self.pre_size = pre_size
self.mnnvl_size = mnnvl_size
self.post_size = post_size
self.pre_padding = pre_padding
self.post_padding = post_padding
self.data_offset = data_offset
self.leading_edge = leading_edge
self.trailing_edge = trailing_edge
self.total_size = total_size
self.handle_phys_size = handle_phys_size
@classmethod
def compute(
cls,
expert_bytes: int,
num_experts: int,
local_start: int,
local_end: int,
granularity: int,
handle_phys_size: int,
pool_granularity: Optional[int] = None,
) -> PageAlignedLayout:
if pool_granularity is None:
pool_granularity = granularity
local_start_bytes = local_start * expert_bytes
local_end_bytes = local_end * expert_bytes
total_expert_bytes = num_experts * expert_bytes
page_start = align_down(local_start_bytes, granularity)
page_end = align_up(local_end_bytes, granularity)
data_offset = local_start_bytes - page_start
leading_edge = data_offset
trailing_edge = page_end - local_end_bytes
mnnvl_size = page_end - page_start
if mnnvl_size > handle_phys_size:
raise ValueError(
f"mnnvl_size ({mnnvl_size}) exceeds handle_phys_size ({handle_phys_size})"
)
pre_size = align_up(page_start, pool_granularity)
pre_padding = pre_size - page_start
post_size_raw = align_up(total_expert_bytes, granularity) - page_end
post_size = align_up(post_size_raw, pool_granularity)
post_padding = post_size - post_size_raw
total_size = pre_size + mnnvl_size + post_size
return cls(
expert_bytes=expert_bytes,
num_experts=num_experts,
local_start=local_start,
local_end=local_end,
granularity=granularity,
pool_granularity=pool_granularity,
page_start=page_start,
page_end=page_end,
pre_size=pre_size,
mnnvl_size=mnnvl_size,
post_size=post_size,
pre_padding=pre_padding,
post_padding=post_padding,
data_offset=data_offset,
leading_edge=leading_edge,
trailing_edge=trailing_edge,
total_size=total_size,
handle_phys_size=handle_phys_size,
)
def get_edge_info(self) -> EdgeInfo:
return EdgeInfo(
data_offset=self.data_offset,
leading_edge=self.leading_edge,
trailing_edge=self.trailing_edge,
page_start=self.page_start,
page_end=self.page_end,
expert_bytes=self.expert_bytes,
)
@property
def pre_pages(self) -> int:
return self.pre_size // self.pool_granularity if self.pool_granularity else 0
@property
def post_pages(self) -> int:
return self.post_size // self.pool_granularity if self.pool_granularity else 0
@property
def remote_pages(self) -> int:
return self.pre_pages + self.post_pages
class MnnvlHandleSet:
def __init__(
self,
handles: Dict[Tuple[int, str], int],
sizes: Dict[Tuple[int, str], int],
):
self.handles = handles
self.sizes = sizes
def get_handle(self, layer_idx: int, name: str) -> int:
return self.handles[(layer_idx, name)]
def get_size(self, layer_idx: int, name: str) -> int:
return self.sizes[(layer_idx, name)]
@property
def layer_indices(self) -> List[int]:
return sorted(set(li for li, _ in self.handles.keys()))
def weight_names(self, layer_idx: int) -> List[str]:
return [n for (li, n) in self.handles.keys() if li == layer_idx]
def compute_peer_ranges(
*,
dwdp_size: int,
num_experts_per_worker: int,
num_prefetch_experts: int,
num_experts_total: int,
) -> PeerRanges:
ranges: PeerRanges = []
for peer_rank in range(dwdp_size):
start = peer_rank * num_prefetch_experts
end_capped = min(start + num_experts_per_worker, num_experts_total)
ranges.append((start, end_capped))
return ranges
def lookup_owner(expert_id: int, peer_ranges: PeerRanges) -> int:
for peer_rank, (start, end) in enumerate(peer_ranges):
if start <= expert_id < end:
return peer_rank
raise ValueError(
f"expert_id={expert_id} not owned by any peer in peer_ranges={peer_ranges}"
)
def build_layer_weight_specs(
local_params: Dict[Tuple[int, str], torch.Tensor],
num_routed_experts: int,
) -> LayerWeightSpecs:
specs: LayerWeightSpecs = {}
for (layer_idx, name), param in local_params.items():
chunk_shape = tuple(param.shape)
specs.setdefault(layer_idx, {})[name] = WeightSpec(
num_experts=num_routed_experts,
chunk_shape=chunk_shape,
full_shape=(num_routed_experts,) + chunk_shape[1:],
dtype=param.dtype,
)
return specs
@@ -0,0 +1,116 @@
# Adapted from NVIDIA TensorRT-LLM (https://github.com/NVIDIA/TensorRT-LLM)
"""Double-buffered pool of local VMM pages backing the remote regions of the composite VA."""
from __future__ import annotations
import logging
from typing import Dict, List, Optional, Tuple
from sglang.srt.layers.moe.dwdp.vmm import (
align_up,
create_local_handle,
get_allocation_granularity,
map_handle,
release_handle,
)
logger = logging.getLogger(__name__)
DEFAULT_PAGE_SIZE_MULTIPLIER = 8
class PagePool:
# local (non-fabric) handles avoid consuming NVLink routing table entries
DEFAULT_PAGE_SIZE_MULTIPLIER = DEFAULT_PAGE_SIZE_MULTIPLIER
def __init__(
self,
slot_sizes: List[int],
device_id: int,
granularity: Optional[int] = None,
page_size: Optional[int] = None,
):
self._device_id = device_id
self._granularity = granularity or get_allocation_granularity(device_id)
if page_size is None:
self._page_size = self.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
else:
self._page_size = page_size
self._slot_sizes = list(slot_sizes)
self._slot_pages = [
align_up(sz, self._page_size) // self._page_size for sz in slot_sizes
]
self._page_handles: List[List[int]] = []
self._released = False
for slot_idx, num_pages in enumerate(self._slot_pages):
handles = []
for _ in range(num_pages):
h = create_local_handle(self._page_size, device_id)
handles.append(h)
self._page_handles.append(handles)
logger.debug(
f"PagePool slot {slot_idx}: {num_pages} pages × {self._page_size} B"
)
@classmethod
def create(
cls,
slot_sizes: List[int],
device_id: int,
page_size: Optional[int] = None,
) -> PagePool:
return cls(slot_sizes, device_id, page_size=page_size)
@property
def page_size(self) -> int:
return self._page_size
def num_pages(self, slot: int) -> int:
return self._slot_pages[slot]
def slot_size(self, slot: int) -> int:
return self._slot_sizes[slot]
def map_pages(
self,
slot: int,
va_start: int,
size: int,
page_offset: int = 0,
) -> List[Tuple[int, int]]:
# does NOT call set_access; caller must set access on the whole composite VA
aligned_size = align_up(size, self._page_size)
num_pages_needed = aligned_size // self._page_size
mappings = []
for i in range(num_pages_needed):
va = va_start + i * self._page_size
handle = self._page_handles[slot][page_offset + i]
map_handle(va, self._page_size, handle, offset=0)
mappings.append((va, self._page_size))
return mappings
def release(self) -> None:
if self._released:
return
self._released = True
for handles in self._page_handles:
for h in handles:
release_handle(h)
self._page_handles = [[], []]
def compute_slot_sizes(
layouts: Dict[int, Dict[str, PageAlignedLayout]], # noqa: F821
buffer_slot_assignments: Dict[int, int],
) -> List[int]:
slot_sizes = [0, 0]
for layer_idx, weight_layouts in layouts.items():
slot = buffer_slot_assignments.get(layer_idx, layer_idx % 2)
total = sum(lo.pre_size + lo.post_size for lo in weight_layouts.values())
slot_sizes[slot] = max(slot_sizes[slot], total)
return slot_sizes
@@ -0,0 +1,238 @@
# Adapted from NVIDIA TensorRT-LLM (https://github.com/NVIDIA/TensorRT-LLM)
"""Cross-rank expert weight handle exchange (FABRIC or POSIX fd) and peer view import."""
from __future__ import annotations
import logging
import os
from typing import Dict, List, Optional, Tuple
import torch
import torch.distributed as dist
from cuda.bindings import driver as cuda
from sglang.srt.distributed.device_communicators.vmm_utils import (
check_drv,
exchange_posix_fds,
export_shareable_handles,
import_peer_handle,
)
from sglang.srt.layers.moe.dwdp.layout import (
DwdpExpertLayout,
LayerWeightSpecs,
MnnvlHandleSet,
)
from sglang.srt.layers.moe.dwdp.vmm import (
align_down,
align_up,
create_fabric_handle,
free_va,
get_allocation_granularity,
map_handle,
release_handle,
reserve_va,
set_access,
tensor_from_ptr,
unmap_va,
)
logger = logging.getLogger(__name__)
def _close_fds(fds) -> None:
for fd in fds:
try:
os.close(fd)
except OSError:
pass
def _copy_local_weights_to_handles(
sorted_keys: List[Tuple[int, str]],
local_params: Dict[Tuple[int, str], torch.Tensor],
layer_weight_specs: LayerWeightSpecs,
layout: DwdpExpertLayout,
device_id: int,
) -> Tuple[Dict[Tuple[int, str], int], Dict[Tuple[int, str], int]]:
granularity = get_allocation_granularity(device_id)
handles: Dict[Tuple[int, str], int] = {}
sizes: Dict[Tuple[int, str], int] = {}
for layer_idx, name in sorted_keys:
param = local_params[(layer_idx, name)]
spec = layer_weight_specs[layer_idx][name]
local_start_bytes = layout.local_expert_start * spec.expert_bytes
local_end_bytes = layout.local_expert_end * spec.expert_bytes
page_start = align_down(local_start_bytes, granularity)
page_end = align_up(local_end_bytes, granularity)
phys_size = page_end - page_start
data_offset = local_start_bytes - page_start
handle = create_fabric_handle(phys_size, device_id)
temp_va = reserve_va(phys_size, granularity)
map_handle(temp_va, phys_size, handle)
set_access(temp_va, phys_size, device_id)
nbytes = param.numel() * param.element_size()
check_drv(
cuda.cuMemcpyDtoD(temp_va + data_offset, param.data_ptr(), nbytes),
"cuMemcpyDtoD",
)
torch.cuda.synchronize()
unmap_va(temp_va, phys_size)
free_va(temp_va, phys_size)
param.untyped_storage().resize_(0)
handles[(layer_idx, name)] = handle
sizes[(layer_idx, name)] = phys_size
logger.debug(
f"Phase 1: layer={layer_idx}, name={name}, "
f"phys_size={phys_size}, data_offset={data_offset}"
)
torch.cuda.empty_cache()
return handles, sizes
class DWDPTransport:
def __init__(self):
self._handle_set: Optional[MnnvlHandleSet] = None
self._peer_views: Dict[Tuple[int, int, str], torch.Tensor] = {}
self._imported_handles: List[int] = []
self._peer_va_regions: List[Tuple[int, int]] = []
@classmethod
def create(
cls,
layer_weight_specs: LayerWeightSpecs,
local_params: Dict[Tuple[int, str], torch.Tensor],
group: dist.ProcessGroup,
layout: DwdpExpertLayout,
device_id: int,
) -> DWDPTransport:
transport = cls()
sorted_keys = sorted(local_params.keys())
handles, sizes = _copy_local_weights_to_handles(
sorted_keys, local_params, layer_weight_specs, layout, device_id
)
transport._handle_set = MnnvlHandleSet(handles=handles, sizes=sizes)
transport._import_peer_views(
sorted_keys, layer_weight_specs, group, layout, device_id
)
dist.barrier(group=group.device_group)
logger.debug(
f"Transport complete: rank={layout.dwdp_rank}/{layout.dwdp_size}, "
f"{len(handles)} handles, {len(transport._peer_views)} peer views"
)
return transport
def _import_peer_views(
self,
sorted_keys: List[Tuple[int, str]],
layer_weight_specs: LayerWeightSpecs,
group: dist.ProcessGroup,
layout: DwdpExpertLayout,
device_id: int,
) -> None:
cpu_group = group.cpu_group
granularity = get_allocation_granularity(device_id)
handle_list = [self._handle_set.get_handle(li, n) for li, n in sorted_keys]
fabric_handles, local_posix_fds, use_fabric = export_shareable_handles(
handle_list, cpu_group, layout.dwdp_rank
)
peer_fds: Dict[Tuple[int, int], int] = {}
key_counts = [None] * layout.dwdp_size
dist.all_gather_object(key_counts, len(sorted_keys), group=cpu_group)
if any(count != len(sorted_keys) for count in key_counts):
raise RuntimeError(
f"Mismatched DWDP weight handle counts across ranks: {key_counts}"
)
if use_fabric:
all_fabric = [None] * layout.dwdp_size
dist.all_gather_object(all_fabric, fabric_handles, group=cpu_group)
else:
all_fabric = None
peer_fds = exchange_posix_fds(
cpu_group,
layout.dwdp_rank,
layout.dwdp_size,
local_posix_fds,
key_counts,
)
logger.info(
"DWDP handle exchange via %s (%d handles)",
"FABRIC" if use_fabric else "POSIX fd",
len(sorted_keys),
)
for key_idx, (layer_idx, name) in enumerate(sorted_keys):
spec = layer_weight_specs[layer_idx][name]
for peer_rank in range(layout.dwdp_size):
if peer_rank == layout.dwdp_rank:
continue
fabric_handle = all_fabric[peer_rank][key_idx] if use_fabric else None
fd = None if use_fabric else peer_fds[(peer_rank, key_idx)]
peer_handle = import_peer_handle(
fabric_handle, fd, use_fabric=use_fabric, peer_rank=peer_rank
)
self._imported_handles.append(int(peer_handle))
peer_start, peer_end = layout.peer_ranges[peer_rank]
peer_start_bytes = peer_start * spec.expert_bytes
peer_end_bytes = peer_end * spec.expert_bytes
peer_page_start = align_down(peer_start_bytes, granularity)
peer_page_end = align_up(peer_end_bytes, granularity)
peer_phys_size = peer_page_end - peer_page_start
peer_data_offset = peer_start_bytes - peer_page_start
peer_va = reserve_va(peer_phys_size, granularity)
map_handle(peer_va, peer_phys_size, int(peer_handle))
set_access(peer_va, peer_phys_size, device_id)
self._peer_va_regions.append((peer_va, peer_phys_size))
num_peer_experts = peer_end - peer_start
peer_tensor = tensor_from_ptr(
ptr=peer_va + peer_data_offset,
shape=(num_peer_experts,) + spec.full_shape[1:],
dtype=spec.dtype,
device_id=device_id,
)
self._peer_views[(peer_rank, layer_idx, name)] = peer_tensor
_close_fds(local_posix_fds)
_close_fds(peer_fds.values())
@property
def handle_set(self) -> MnnvlHandleSet:
assert self._handle_set is not None
return self._handle_set
@property
def peer_views(self) -> Dict[Tuple[int, int, str], torch.Tensor]:
return self._peer_views
def release(self) -> None:
for va, size in self._peer_va_regions:
unmap_va(va, size)
free_va(va, size)
self._peer_va_regions.clear()
for h in self._imported_handles:
release_handle(h)
self._imported_handles.clear()
self._peer_views.clear()
+258
View File
@@ -0,0 +1,258 @@
"""CUDA VMM primitives for DWDP: handle creation, VA reserve/map, DLPack tensor views."""
from __future__ import annotations
import ctypes
import functools
import logging
from typing import Tuple
import torch
from cuda.bindings import driver as cuda
from sglang.srt.distributed.device_communicators.vmm_utils import (
check_drv,
make_rw_access_desc,
)
logger = logging.getLogger(__name__)
def align_up(value: int, alignment: int) -> int:
if alignment <= 0 or (alignment & (alignment - 1)) != 0:
raise ValueError(f"alignment must be a positive power of 2, got {alignment}")
return ((value + alignment - 1) // alignment) * alignment
def align_down(value: int, alignment: int) -> int:
if alignment <= 0 or (alignment & (alignment - 1)) != 0:
raise ValueError(f"alignment must be a positive power of 2, got {alignment}")
return (value // alignment) * alignment
def _make_prop(device_id: int, handle_types: int) -> cuda.CUmemAllocationProp:
prop = cuda.CUmemAllocationProp()
prop.type = cuda.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
prop.location.type = cuda.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
prop.location.id = device_id
prop.requestedHandleTypes = handle_types
return prop
@functools.lru_cache(maxsize=None)
def shareable_handle_types(device_id: int) -> int:
fabric = int(cuda.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC)
posix = int(cuda.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR)
fabric_supported = check_drv(
cuda.cuDeviceGetAttribute(
cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_HANDLE_TYPE_FABRIC_SUPPORTED,
device_id,
),
"cuDeviceGetAttribute(FABRIC_SUPPORTED)",
)
if fabric_supported:
# the attribute alone is not sufficient: drivers advertise FABRIC on
# platforms where creation still fails (e.g. no IMEX channel), so a
# real cuMemCreate probe decides
combined = fabric | posix
option = (
cuda.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED
)
try:
prop = _make_prop(device_id, combined)
gran = check_drv(
cuda.cuMemGetAllocationGranularity(prop=prop, option=option),
"cuMemGetAllocationGranularity(probe)",
)
handle = check_drv(
cuda.cuMemCreate(int(gran), prop, 0), "cuMemCreate(probe)"
)
check_drv(cuda.cuMemRelease(handle), "cuMemRelease(probe)")
return combined
except RuntimeError as e:
logger.info(
"FABRIC advertised on device %s but creation probe failed (%s); "
"DWDP handles will be POSIX fd only",
device_id,
e,
)
return posix
@functools.lru_cache(maxsize=None)
def get_allocation_granularity(device_id: int) -> int:
prop = _make_prop(device_id, shareable_handle_types(device_id))
option = cuda.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED
return check_drv(
cuda.cuMemGetAllocationGranularity(prop=prop, option=option),
"cuMemGetAllocationGranularity",
)
def create_fabric_handle(size: int, device_id: int) -> int:
prop = _make_prop(device_id, shareable_handle_types(device_id))
handle = check_drv(cuda.cuMemCreate(size, prop, flags=0), "cuMemCreate")
return int(handle)
def create_local_handle(size: int, device_id: int) -> int:
# non-shareable handle: does not consume a fabric routing table entry
prop = _make_prop(device_id, 0)
handle = check_drv(cuda.cuMemCreate(size, prop, flags=0), "cuMemCreate(local)")
return int(handle)
def release_handle(handle: int) -> None:
if handle != 0:
check_drv(cuda.cuMemRelease(handle), "cuMemRelease")
def reserve_va(size: int, granularity: int) -> int:
va = check_drv(
cuda.cuMemAddressReserve(size, granularity, 0, 0), "cuMemAddressReserve"
)
return int(va)
def free_va(va: int, size: int) -> None:
if va != 0:
check_drv(cuda.cuMemAddressFree(va, size), "cuMemAddressFree")
def map_handle(va: int, size: int, handle: int, offset: int = 0) -> None:
check_drv(cuda.cuMemMap(va, size, offset, handle, 0), "cuMemMap")
def unmap_va(va: int, size: int) -> None:
check_drv(cuda.cuMemUnmap(va, size), "cuMemUnmap")
def set_access(va: int, size: int, device_id: int) -> None:
desc = make_rw_access_desc(device_id)
check_drv(cuda.cuMemSetAccess(va, size, [desc], 1), "cuMemSetAccess")
class _DLDataType(ctypes.Structure):
_fields_ = [
("code", ctypes.c_uint8),
("bits", ctypes.c_uint8),
("lanes", ctypes.c_uint16),
]
class _DLDevice(ctypes.Structure):
_fields_ = [("device_type", ctypes.c_int), ("device_id", ctypes.c_int)]
class _DLTensor(ctypes.Structure):
_fields_ = [
("data", ctypes.c_void_p),
("device", _DLDevice),
("ndim", ctypes.c_int),
("dtype", _DLDataType),
("shape", ctypes.POINTER(ctypes.c_int64)),
("strides", ctypes.POINTER(ctypes.c_int64)),
("byte_offset", ctypes.c_size_t),
]
class _DLManagedTensor(ctypes.Structure):
pass
_DLManagedTensor._fields_ = [
("dl_tensor", _DLTensor),
("manager_ctx", ctypes.c_void_p),
("deleter", ctypes.CFUNCTYPE(None, ctypes.POINTER(_DLManagedTensor))),
]
@ctypes.CFUNCTYPE(None, ctypes.POINTER(_DLManagedTensor))
def _no_op_deleter(_ptr):
pass
_FLOAT8_DTYPES = {
torch.float8_e5m2,
torch.float8_e4m3fn,
torch.float8_e4m3fnuz,
torch.float8_e5m2fnuz,
}
def _torch_dtype_to_dl(dtype: torch.dtype) -> Tuple[int, int]:
# float8 goes through DLPack as uint8 (from_dlpack rejects kFloat/8-bit); caller view-casts back
if dtype in _FLOAT8_DTYPES:
return 1, 8
if dtype in (
torch.bfloat16,
torch.float16,
torch.float32,
torch.float64,
):
return 2, torch.finfo(dtype).bits
if dtype in (torch.int8, torch.int16, torch.int32, torch.int64):
return 0, torch.iinfo(dtype).bits
if dtype in (torch.uint8,):
return 1, 8
raise NotImplementedError(f"Unsupported dtype for DLPack: {dtype}")
def tensor_from_ptr(
ptr: int,
shape: Tuple[int, ...],
dtype: torch.dtype,
device_id: int,
) -> torch.Tensor:
if ptr == 0:
raise ValueError("Cannot create tensor from null pointer")
numel = 1
for d in shape:
if d <= 0:
raise ValueError(f"All dimensions must be positive, got shape={shape}")
numel *= d
dl_code, bits = _torch_dtype_to_dl(dtype)
ndim = len(shape)
ShapeArray = ctypes.c_int64 * ndim
shape_arr = ShapeArray(*shape)
device = _DLDevice(device_type=2, device_id=device_id) # kDLCUDA = 2
dl_dtype = _DLDataType(code=dl_code, bits=bits, lanes=1)
dl_tensor = _DLTensor()
dl_tensor.data = ctypes.c_void_p(ptr)
dl_tensor.device = device
dl_tensor.ndim = ndim
dl_tensor.dtype = dl_dtype
dl_tensor.shape = ctypes.cast(shape_arr, ctypes.POINTER(ctypes.c_int64))
dl_tensor.strides = None
dl_tensor.byte_offset = 0
managed = _DLManagedTensor()
managed.dl_tensor = dl_tensor
managed.manager_ctx = None
managed.deleter = _no_op_deleter
ctypes.pythonapi.PyCapsule_New.restype = ctypes.c_void_p
ctypes.pythonapi.PyCapsule_New.argtypes = [
ctypes.c_void_p,
ctypes.c_char_p,
ctypes.c_void_p,
]
capsule_ptr = ctypes.pythonapi.PyCapsule_New(
ctypes.pointer(managed),
b"dltensor",
None,
)
capsule = ctypes.cast(capsule_ptr, ctypes.py_object).value
tensor = torch.utils.dlpack.from_dlpack(capsule)
tensor = tensor.reshape(shape)
if dtype in _FLOAT8_DTYPES:
tensor = tensor.view(dtype)
# ctypes structs must outlive the tensor or the data pointer dangles
tensor._dlpack_prevent_gc = (shape_arr, managed, capsule)
return tensor
@@ -0,0 +1,221 @@
# Adapted from NVIDIA TensorRT-LLM (https://github.com/NVIDIA/TensorRT-LLM)
"""Composite VA presenting a contiguous full-expert weight tensor per (layer, weight)."""
from __future__ import annotations
import logging
from typing import Dict, List, Optional, Tuple
import torch
from sglang.srt.layers.moe.dwdp.layout import (
EdgeInfo,
LayerWeightSpecs,
MnnvlHandleSet,
PageAlignedLayout,
)
from sglang.srt.layers.moe.dwdp.page_pool import PagePool, compute_slot_sizes
from sglang.srt.layers.moe.dwdp.vmm import (
free_va,
get_allocation_granularity,
map_handle,
reserve_va,
set_access,
tensor_from_ptr,
unmap_va,
)
logger = logging.getLogger(__name__)
class WeightBuffer:
def __init__(
self,
layer_weight_specs: LayerWeightSpecs,
handles: MnnvlHandleSet,
local_start: int,
local_end: int,
dwdp_size: int,
device_id: int,
):
self._layer_weight_specs = layer_weight_specs
self._handles = handles
self._local_start = local_start
self._local_end = local_end
self._dwdp_size = dwdp_size
self._device_id = device_id
self._granularity = get_allocation_granularity(device_id)
self._pool_page_size = PagePool.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
self._page_pool: Optional[PagePool] = None
self._moe_layer_indices = sorted(layer_weight_specs.keys())
self._layouts: Dict[int, Dict[str, PageAlignedLayout]] = {}
self._tensors: Dict[int, Dict[str, torch.Tensor]] = {}
self._remote_slices: Dict[
int, Dict[str, List[Tuple[torch.Tensor, int, int]]]
] = {}
self._mappings: Dict[int, List[Tuple[int, int]]] = {}
self._va_regions: Dict[int, List[Tuple[int, int]]] = {}
self._released = False
@classmethod
def create(
cls,
layer_weight_specs: LayerWeightSpecs,
handles: MnnvlHandleSet,
local_start: int,
local_end: int,
dwdp_size: int,
device_id: int,
) -> WeightBuffer:
buf = cls(
layer_weight_specs, handles, local_start, local_end, dwdp_size, device_id
)
for li, ws in layer_weight_specs.items():
buf._layouts[li] = {}
for name, spec in ws.items():
buf._layouts[li][name] = PageAlignedLayout.compute(
expert_bytes=spec.expert_bytes,
num_experts=spec.num_experts,
local_start=local_start,
local_end=local_end,
granularity=buf._granularity,
handle_phys_size=handles.get_size(li, name),
pool_granularity=buf._pool_page_size,
)
assignments = {li: buf.buffer_index_for_layer(li) for li in layer_weight_specs}
slot_sizes = compute_slot_sizes(buf._layouts, assignments)
buf._page_pool = PagePool.create(
slot_sizes, device_id, page_size=buf._pool_page_size
)
for li in buf._moe_layer_indices:
buf._setup_layer(li)
logger.debug(
f"WeightBuffer created for {len(buf._moe_layer_indices)} layers, "
f"local [{local_start}, {local_end})"
)
return buf
def _setup_layer(self, layer_idx: int) -> None:
weight_layouts = self._layouts[layer_idx]
weight_specs = self._layer_weight_specs[layer_idx]
buf_slot = self.buffer_index_for_layer(layer_idx)
self._tensors[layer_idx] = {}
self._remote_slices[layer_idx] = {}
self._mappings[layer_idx] = []
self._va_regions[layer_idx] = []
page_pool_offset = 0
for name, layout in weight_layouts.items():
spec = weight_specs[name]
handle = self._handles.get_handle(layer_idx, name)
va_base = reserve_va(layout.total_size, self._granularity)
self._va_regions[layer_idx].append((va_base, layout.total_size))
all_maps = self._mappings[layer_idx]
if layout.pre_size > 0:
pre_maps = self._page_pool.map_pages(
slot=buf_slot,
va_start=va_base,
size=layout.pre_size,
page_offset=page_pool_offset,
)
all_maps.extend(pre_maps)
page_pool_offset += layout.pre_pages
mnnvl_va = va_base + layout.pre_size
map_handle(mnnvl_va, layout.mnnvl_size, handle, offset=0)
all_maps.append((mnnvl_va, layout.mnnvl_size))
if layout.post_size > 0:
post_va = mnnvl_va + layout.mnnvl_size
post_maps = self._page_pool.map_pages(
slot=buf_slot,
va_start=post_va,
size=layout.post_size,
page_offset=page_pool_offset,
)
all_maps.extend(post_maps)
page_pool_offset += layout.post_pages
set_access(va_base, layout.total_size, self._device_id)
tensor_start = va_base + layout.pre_padding
full_tensor = tensor_from_ptr(
ptr=tensor_start,
shape=spec.full_shape,
dtype=spec.dtype,
device_id=self._device_id,
)
self._tensors[layer_idx][name] = full_tensor
slices = []
if self._local_start > 0:
slices.append((full_tensor[: self._local_start], 0, self._local_start))
if self._local_end < spec.num_experts:
slices.append(
(full_tensor[self._local_end :], self._local_end, spec.num_experts)
)
self._remote_slices[layer_idx][name] = slices
def get_full_tensor(self, layer_idx: int, name: str) -> torch.Tensor:
return self._tensors[layer_idx][name]
def get_remote_slices(
self, layer_idx: int, name: str
) -> List[Tuple[torch.Tensor, int, int]]:
return self._remote_slices[layer_idx][name]
def get_edge_info(self, layer_idx: int, name: str) -> EdgeInfo:
return self._layouts[layer_idx][name].get_edge_info()
def get_layout(self, layer_idx: int, name: str) -> PageAlignedLayout:
return self._layouts[layer_idx][name]
@property
def layer_indices(self) -> List[int]:
return list(self._moe_layer_indices)
@property
def local_start(self) -> int:
return self._local_start
@property
def local_end(self) -> int:
return self._local_end
@property
def device_id(self) -> int:
return self._device_id
def weight_names(self, layer_idx: int) -> List[str]:
return list(self._layer_weight_specs[layer_idx].keys())
def buffer_index_for_layer(self, layer_idx: int) -> int:
if layer_idx in self._moe_layer_indices:
return self._moe_layer_indices.index(layer_idx) % 2
return layer_idx % 2
def release(self) -> None:
if self._released:
return
self._released = True
for li, maps in self._mappings.items():
for va, sz in maps:
unmap_va(va, sz)
for li, regions in self._va_regions.items():
for va, sz in regions:
free_va(va, sz)
self._mappings.clear()
self._va_regions.clear()
self._tensors.clear()
self._remote_slices.clear()
if self._page_pool is not None:
self._page_pool.release()
self._page_pool = None
@@ -0,0 +1,142 @@
# Adapted from NVIDIA TensorRT-LLM (https://github.com/NVIDIA/TensorRT-LLM)
"""Double-buffered async prefetch of peer expert weights into the composite VA."""
from __future__ import annotations
import bisect
import logging
from typing import Dict, List, Optional, Tuple
import torch
from sglang.srt.layers.moe.dwdp.layout import PeerRanges, lookup_owner
from sglang.srt.layers.moe.dwdp.weight_buffer import WeightBuffer
logger = logging.getLogger(__name__)
class DWDPWeightManager:
def __init__(
self,
weight_buffer: WeightBuffer,
peer_views: Dict[Tuple[int, int, str], torch.Tensor],
peer_ranges: PeerRanges,
moe_layer_indices: List[int],
weight_names: List[str],
dwdp_rank: int,
dwdp_size: int,
transport=None,
) -> None:
self._weight_buffer = weight_buffer
self._peer_views = peer_views
self._peer_ranges = peer_ranges
self._moe_layer_indices = sorted(moe_layer_indices)
self._moe_layer_set = set(self._moe_layer_indices)
self._weight_names = list(weight_names)
self._dwdp_rank = dwdp_rank
self._dwdp_size = dwdp_size
# transport handles underpin the VA mappings; must outlive this manager
self._transport = transport
device = torch.device("cuda", weight_buffer.device_id)
self._copy_stream = torch.cuda.Stream(device=device)
self._prefetch_events: List[torch.cuda.Event] = [
torch.cuda.Event() for _ in range(2)
]
self._consume_events: List[torch.cuda.Event] = [
torch.cuda.Event() for _ in range(2)
]
# pre-record consume events so the first prefetch doesn't stall
current = torch.cuda.current_stream(device)
for ev in self._consume_events:
ev.record(current)
logger.debug(
f"WeightManager rank={dwdp_rank}/{dwdp_size}, "
f"{len(moe_layer_indices)} MoE layers, weights={weight_names}"
)
@property
def weight_buffer(self) -> WeightBuffer:
return self._weight_buffer
def is_moe_layer(self, layer_idx: int) -> bool:
return layer_idx in self._moe_layer_set
def next_moe_layer(self, layer_idx: int) -> Optional[int]:
pos = bisect.bisect_right(self._moe_layer_indices, layer_idx)
if pos < len(self._moe_layer_indices):
return self._moe_layer_indices[pos]
return None
def first_moe_layer(self) -> int:
return self._moe_layer_indices[0]
def prefetch_layer(self, layer_idx: int) -> None:
buf_idx = self._weight_buffer.buffer_index_for_layer(layer_idx)
with torch.cuda.stream(self._copy_stream):
# WAR: wait for compute to finish reading this slot before overwriting
self._copy_stream.wait_event(self._consume_events[buf_idx])
self._prefetch_layer_per_slice(layer_idx)
self._prefetch_events[buf_idx].record(self._copy_stream)
def _prefetch_layer_per_slice(self, layer_idx: int) -> None:
for name in self._weight_names:
remote_slices = self._weight_buffer.get_remote_slices(layer_idx, name)
for dst_slice, expert_start, expert_end in remote_slices:
cursor = expert_start
dst_offset = 0
while cursor < expert_end:
peer_rank = lookup_owner(cursor, self._peer_ranges)
peer_start, peer_end = self._peer_ranges[peer_rank]
local_offset = cursor - peer_start
chunk_end = min(expert_end, peer_end)
n = chunk_end - cursor
peer_key = (peer_rank, layer_idx, name)
src = self._peer_views[peer_key]
dst_slice[dst_offset : dst_offset + n].copy_(
src[local_offset : local_offset + n]
)
dst_offset += n
cursor = chunk_end
def wait_prefetch(self, layer_idx: int) -> None:
buf_idx = self._weight_buffer.buffer_index_for_layer(layer_idx)
device = torch.device("cuda", self._weight_buffer.device_id)
compute_stream = torch.cuda.current_stream(device)
compute_stream.wait_event(self._prefetch_events[buf_idx])
def record_compute_and_prefetch_next(self, layer_idx: int) -> None:
buf_idx = self._weight_buffer.buffer_index_for_layer(layer_idx)
device = torch.device("cuda", self._weight_buffer.device_id)
compute_stream = torch.cuda.current_stream(device)
self._consume_events[buf_idx].record(compute_stream)
# prefetch the layer 2 ahead — it reuses the same buffer slot
next_layer = self.next_moe_layer(layer_idx)
if next_layer is not None:
next_next = self.next_moe_layer(next_layer)
if next_next is not None:
self.prefetch_layer(next_next)
def prefetch_first_layers(self) -> None:
if len(self._moe_layer_indices) >= 1:
self.prefetch_layer(self._moe_layer_indices[0])
if len(self._moe_layer_indices) >= 2:
self.prefetch_layer(self._moe_layer_indices[1])
def release(self) -> None:
if self._weight_buffer is not None:
self._weight_buffer.release()
self._weight_buffer = None
if self._transport is not None:
self._transport.release()
self._transport = None
self._peer_views.clear()
@@ -5,7 +5,7 @@
import logging
from enum import Enum
from functools import cached_property
from typing import List, Optional, Tuple
from typing import Dict, List, Optional, Tuple
import torch
from torch.nn.parameter import UninitializedParameter
@@ -70,7 +70,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
is_in_tc_piecewise_cuda_graph,
)
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.runtime_context import (
get_global_dwdp_manager,
get_parallel,
get_server_args,
)
from sglang.srt.utils import (
cpu_has_amx_support,
get_bool_env_var,
@@ -389,9 +393,71 @@ class FusedMoE(torch.nn.Module):
self.down_gemm_overlap_args: Optional[DownGemmOverlapArgs] = None
self.meta_overlap_args: Optional[dict] = None
self._dwdp_bound = False
if self.quant_method is not None and hasattr(self.quant_method, "runner"):
self.runner = self.quant_method.runner
@property
def num_global_routed_experts(self) -> int:
return self._num_global_routed
def bind_full_expert_weights(self, weights: Dict[str, torch.Tensor]) -> None:
"""Rebind this layer's expert weight tensors to externally provided
full [num_experts, ...] tensors and collapse its EP view to a single
rank that owns every routed expert (ep_size=1, no expert remapping).
Callers are weight-replication schemes that materialize all expert
weights locally after load time (e.g. DWDP's composite-VA prefetch).
"""
self.moe_ep_size = 1
self.moe_ep_rank = 0
self._num_local_routed = self._num_global_routed
self.num_local_experts = self.num_experts
self.moe_runner_config.num_local_experts = self.num_local_experts
self.dispatcher.moe_ep_size = 1
self.dispatcher.moe_ep_rank = 0
self.dispatcher.num_local_experts = self.num_local_experts
self.dispatcher.num_local_routed_experts = self._num_local_routed
self.dispatcher.local_expert_mapping = None
self.dispatcher.expert_mask_gpu = None
for name, tensor in weights.items():
self.replace_expert_tensor(name, tensor)
self._dwdp_bound = True
def named_per_expert_tensors(
self, num_local_experts: int
) -> List[Tuple[str, torch.Tensor]]:
"""Expert-sharded side tensors (dim0 == num_local_experts): quant
scales, alphas, biases. Excludes the main w13/w2 expert weights."""
found: Dict[str, torch.Tensor] = {}
for name, param in self._parameters.items():
if param is not None:
found[name] = param.data
for name, buf in self._buffers.items():
if buf is not None and name not in found:
found[name] = buf
for name, value in vars(self).items():
if isinstance(value, torch.Tensor) and name not in found:
found[name] = value
return [
(name, tensor)
for name, tensor in sorted(found.items())
if name not in ("w13_weight", "w2_weight")
and tensor.ndim > 0
and tensor.shape[0] == num_local_experts
]
def replace_expert_tensor(self, name: str, tensor: torch.Tensor) -> None:
param = self._parameters.get(name)
if param is not None:
param.data = tensor
else:
setattr(self, name, tensor)
@cached_property
def use_padded_loading(self) -> bool:
# This handles the case where the loaded weights are smaller than the padded expert_data
@@ -1278,6 +1344,10 @@ class FusedMoE(torch.nn.Module):
origin_hidden_states_dim = hidden_states.shape[-1]
assert self.quant_method is not None
if self._dwdp_bound:
dwdp_mgr = get_global_dwdp_manager()
dwdp_mgr.wait_prefetch(self.layer_id)
dispatch_output = self.dispatcher.dispatch(
hidden_states=hidden_states, topk_output=topk_output
)
@@ -1286,6 +1356,9 @@ class FusedMoE(torch.nn.Module):
dispatch_output=dispatch_output,
)
if self._dwdp_bound:
dwdp_mgr.record_compute_and_prefetch_next(self.layer_id)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
+2
View File
@@ -490,6 +490,8 @@ def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool:
"""
if should_skip_mlp_all_reduce():
return True
if get_server_args().dwdp_size > 1:
return True
if should_use_dp_reduce_scatterv():
return True
if is_tp_path and should_use_flashinfer_cutlass_moe_fp4_allgather():
@@ -226,6 +226,7 @@ def prepare_mlp_sync_batch_raw(
require_mlp_tp_gather: bool,
disable_overlap_schedule: bool,
offload_tags: set[str],
dwdp: bool = False,
):
# Check if other DP workers have running batches
if (
@@ -327,8 +328,8 @@ def prepare_mlp_sync_batch_raw(
# Decide whether to emit idle batch
if skip_all_gather:
# Skip idle batch when attn-dp=1
need_idle_batch = dp_size > 1
# Skip idle batch when attn-dp=1 (and always under DWDP: ranks run independently)
need_idle_batch = not dwdp and dp_size > 1
else:
need_idle_batch = max(mlp_sync_info.global_num_tokens) > 0
@@ -386,6 +387,7 @@ class SchedulerDPAttnAdapter:
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
offload_tags=self.offload_tags,
dwdp=self.server_args.dwdp_size > 1,
)
def maybe_prepare_mlp_sync_batch(
@@ -78,6 +78,7 @@ from sglang.srt.layers.cp.utils import (
get_cp_strategy,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.moe.dwdp import DwdpManager
from sglang.srt.layers.sampler import create_sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled
@@ -160,7 +161,11 @@ from sglang.srt.model_executor.runner import (
get_batch_sizes_to_capture,
)
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_server_args
from sglang.srt.runtime_context import (
get_global_dwdp_manager,
get_server_args,
set_global_dwdp_manager,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ( # noqa: F401 (re-export)
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
@@ -572,6 +577,9 @@ class ModelRunner:
moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.ps.moe_ep_rank,
)
self.maybe_init_dwdp()
# Must run before backend/graph init so no draft graph records a
# routed-experts capture-write kernel.
if self.is_draft_worker:
@@ -1014,6 +1022,15 @@ class ModelRunner:
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
)
def maybe_init_dwdp(self):
if self.is_draft_worker:
return
if self.server_args.dwdp_size <= 1:
return
manager = DwdpManager(self.server_args)
set_global_dwdp_manager(manager)
manager.setup(self.model)
def init_lora_manager(self):
self.lora_manager = LoRAManager(
base_model=self.model,
@@ -1424,6 +1441,10 @@ class ModelRunner:
# dispatch below reads the pool.
self._maybe_execute_deferred_mamba_cow_and_clear(forward_batch)
dwdp_mgr = get_global_dwdp_manager()
if dwdp_mgr is not None:
dwdp_mgr.prefetch_first_layers()
if forward_batch.forward_mode.is_split_prefill():
# Layer-split mode; stays on ModelRunner, not the eager runner.
ret = self.forward_split_prefill(
+36 -2
View File
@@ -68,7 +68,11 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args
from sglang.srt.runtime_context import (
get_forward,
get_parallel,
get_server_args,
)
from sglang.srt.utils import (
LazyValue,
add_prefix,
@@ -254,10 +258,40 @@ class GptOssSparseMoeBlock(nn.Module):
hidden_states: torch.Tensor,
forward_batch: Optional[ForwardBatch] = None,
) -> torch.Tensor:
if get_server_args().dwdp_size > 1:
return self.forward_dwdp(hidden_states)
if not get_moe_a2a_backend().is_deepep():
return self.forward_normal(hidden_states)
else:
raise Exception("forward_deepep branch not implemented yet")
raise NotImplementedError("forward_deepep branch not implemented yet")
def forward_dwdp(
self,
hidden_states: torch.Tensor,
) -> torch.Tensor:
num_tokens = hidden_states.shape[0]
hidden_dim_unpadded = self.hidden_size
is_prepadded = hidden_states.shape[-1] != hidden_dim_unpadded
if num_tokens > 0:
router_input = (
hidden_states[..., :hidden_dim_unpadded]
if is_prepadded
else hidden_states
)
router_logits, _ = self.router(router_input)
topk_output = self.topk(router_input, router_logits)
final_hidden_states = self.experts(hidden_states, topk_output)
else:
final_hidden_states = hidden_states
if is_prepadded:
ans = final_hidden_states[..., :hidden_dim_unpadded].contiguous()
ans = ans.view(num_tokens, hidden_dim_unpadded)
else:
ans = final_hidden_states.view(num_tokens, hidden_dim_unpadded)
return ans
def get_moe_weights(self):
return [
+194 -16
View File
@@ -13,6 +13,8 @@
# ==============================================================================
import logging
import math
import re
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
import torch
@@ -76,10 +78,15 @@ from sglang.srt.model_loader.weight_utils import (
)
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args
from sglang.srt.runtime_context import (
get_forward,
get_parallel,
get_server_args,
)
from sglang.srt.utils import (
LazyValue,
add_prefix,
ceil_align,
is_non_idle_and_non_empty,
make_layers,
)
@@ -90,35 +97,193 @@ logger = logging.getLogger(__name__)
def load_mimo_v2_qkv_proj_weight(
name, param, loaded_weight, expected_fused_tp_size: Optional[int] = None
name,
param,
loaded_weight,
expected_fused_tp_size: Optional[int] = None,
deferred_scale_inv: Optional[Dict[str, torch.Tensor]] = None,
):
if loaded_weight.shape == param.shape:
# The checkpoint already stores this rank's qkv_proj shard.
tp_size = get_parallel().attn_tp_size
tp_rank = get_parallel().attn_tp_rank
ckpt_tp = expected_fused_tp_size if expected_fused_tp_size is not None else tp_size
if ckpt_tp == tp_size and loaded_weight.shape == param.shape:
default_weight_loader(param, loaded_weight)
return
if expected_fused_tp_size is not None and expected_fused_tp_size % tp_size != 0:
raise ValueError(
f"MiMoV2 fused qkv_proj checkpoint is TP={expected_fused_tp_size}-"
f"interleaved; got attention tp_size={tp_size} while loading {name}."
)
is_scale_inv = "weight_scale_inv" in name
if is_scale_inv and ckpt_tp != tp_size:
if deferred_scale_inv is not None:
deferred_scale_inv[name] = loaded_weight.clone()
return
raise ValueError(
f"qkv_proj scale_inv {name}: shape mismatch "
f"{tuple(loaded_weight.shape)} vs {tuple(param.shape)} "
f"due to block quantization ceiling; pass deferred_scale_inv dict"
)
if loaded_weight.ndim != param.ndim or loaded_weight.shape[1:] != param.shape[1:]:
raise ValueError(
f"qkv_proj weight {name}: unexpected shape {tuple(loaded_weight.shape)}; "
f"expected sharded {tuple(param.shape)}"
)
if tp_size == ckpt_tp:
fused_shape = (param.shape[0] * tp_size, *param.shape[1:])
if tuple(loaded_weight.shape) != fused_shape:
raise ValueError(
f"qkv_proj weight {name}: unexpected shape "
f"{tuple(loaded_weight.shape)}; expected fused {fused_shape} "
f"or sharded {tuple(param.shape)}"
)
default_weight_loader(param, loaded_weight.chunk(tp_size, dim=0)[tp_rank])
else:
shards_per_rank = ckpt_tp // tp_size
shards = loaded_weight.chunk(ckpt_tp, dim=0)
merged = torch.cat(
shards[tp_rank * shards_per_rank : (tp_rank + 1) * shards_per_rank],
dim=0,
)
default_weight_loader(param, merged)
def _get_ckpt_qkv_shard_sizes(config, layer_name, ckpt_tp):
m = re.search(r"layers\.(\d+)\.", layer_name)
if m is None:
return None
layer_id = int(m.group(1))
pattern = getattr(config, "hybrid_layer_pattern", None)
is_swa = pattern is not None and pattern[layer_id] == 1
if is_swa:
nh = config.swa_num_attention_heads
nkv = config.swa_num_key_value_heads
hd = config.swa_head_dim
vhd = getattr(config, "swa_v_head_dim", hd)
else:
nh = config.num_attention_heads
nkv = config.num_key_value_heads
hd = config.head_dim
vhd = getattr(config, "v_head_dim", hd)
q_per_shard = (nh // ckpt_tp) * hd
k_per_shard = max(1, nkv // ckpt_tp) * hd
v_per_shard = max(1, nkv // ckpt_tp) * vhd
return (q_per_shard, k_per_shard, v_per_shard)
def _deinterleave_qkv_shards(shards, q_per_shard, k_per_shard, v_per_shard):
all_q, all_k, all_v = [], [], []
for s in shards:
all_q.append(s[:q_per_shard])
all_k.append(s[q_per_shard : q_per_shard + k_per_shard])
all_v.append(s[q_per_shard + k_per_shard :])
return torch.cat(all_q + all_k + all_v, dim=0)
def _resolve_deferred_qkv_scale_inv(
params_dict: Dict[str, torch.nn.Parameter],
deferred_scale_inv: Dict[str, torch.Tensor],
expected_fused_tp_size: int,
block_size: int = 128,
config=None,
):
tp_size = get_parallel().attn_tp_size
tp_rank = get_parallel().attn_tp_rank
if expected_fused_tp_size is not None and tp_size != expected_fused_tp_size:
raise ValueError(
f"MiMoV2 fused qkv_proj checkpoint is TP={expected_fused_tp_size}-"
f"interleaved; got attention tp_size={tp_size} while loading {name}."
ckpt_tp = expected_fused_tp_size
shards_per_rank = ckpt_tp // tp_size
for scale_name, ckpt_scale in deferred_scale_inv.items():
weight_name = scale_name.replace(".weight_scale_inv", ".weight")
if weight_name not in params_dict:
raise ValueError(
f"Cannot resolve deferred scale_inv {scale_name}: "
f"weight {weight_name} not found"
)
weight_param = params_dict[weight_name]
scale_param = params_dict[scale_name]
weight_data = weight_param.data
ckpt_scale_shards = ckpt_scale.chunk(ckpt_tp, dim=0)
my_scale_shards = ckpt_scale_shards[
tp_rank * shards_per_rank : (tp_rank + 1) * shards_per_rank
]
weight_rows = weight_data.shape[0]
rows_per_ckpt_shard = weight_rows // shards_per_rank
block_k = ckpt_scale.shape[1]
device = weight_data.device
dequant_shards = []
for i, shard_scale in enumerate(my_scale_shards):
shard_weight = weight_data[
i * rows_per_ckpt_shard : (i + 1) * rows_per_ckpt_shard
]
shard_scale_f32 = shard_scale.to(dtype=torch.float32, device=device)
scale_expanded = shard_scale_f32.repeat_interleave(
block_size, dim=0
).repeat_interleave(block_size, dim=1)
scale_expanded = scale_expanded[
: shard_weight.shape[0], : shard_weight.shape[1]
]
dequant_shards.append(
(shard_weight.to(torch.float32) * scale_expanded).to(torch.bfloat16)
)
qkv_sizes = (
_get_ckpt_qkv_shard_sizes(config, scale_name, ckpt_tp)
if config is not None
else None
)
if qkv_sizes is not None and shards_per_rank > 1:
merged_bf16 = _deinterleave_qkv_shards(dequant_shards, *qkv_sizes)
else:
merged_bf16 = torch.cat(dequant_shards, dim=0)
n, k = merged_bf16.shape
n_pad = ceil_align(n, block_size)
k_pad = ceil_align(k, block_size)
padded = torch.zeros(
n_pad, k_pad, dtype=merged_bf16.dtype, device=merged_bf16.device
)
padded[:n, :k] = merged_bf16
blocks = padded.view(
n_pad // block_size, block_size, k_pad // block_size, block_size
)
amax = blocks.abs().float().amax(dim=(1, 3), keepdim=True).clamp(min=1e-12)
finfo = torch.finfo(torch.float8_e4m3fn)
new_scale_inv = amax / finfo.max
new_fp8 = (
(blocks.float() / new_scale_inv)
.clamp(min=finfo.min, max=finfo.max)
.to(torch.float8_e4m3fn)
)
fused_shape = (param.shape[0] * tp_size, *param.shape[1:])
if tuple(loaded_weight.shape) != fused_shape:
raise ValueError(
f"qkv_proj weight {name}: unexpected shape {tuple(loaded_weight.shape)}; "
f"expected fused {fused_shape} or sharded {tuple(param.shape)}"
n_scale = math.ceil(n / block_size)
k_scale = math.ceil(k / block_size)
weight_param.data.copy_(new_fp8.view(n_pad, k_pad)[:n, :k].contiguous())
default_weight_loader(
scale_param,
new_scale_inv.view(n_pad // block_size, k_pad // block_size)[
:n_scale, :k_scale
].contiguous(),
)
default_weight_loader(param, loaded_weight.chunk(tp_size, dim=0)[tp_rank])
logger.info(
f"Resolved deferred qkv scale_inv {scale_name}: "
f"dequant {shards_per_rank} shards -> requant "
f"({n_scale}, {k_scale})"
+ (f", de-interleaved QKV {qkv_sizes}" if qkv_sizes else "")
)
class MiMoV2MLP(nn.Module):
@@ -323,7 +488,6 @@ class MiMoV2MoE(nn.Module):
) -> torch.Tensor:
if hidden_states.shape[0] > 0:
# router_logits: (num_tokens, n_experts)
router_logits = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
else:
@@ -1257,6 +1421,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
params_dict = dict(self.named_parameters())
skipped_mtp_weights = False
deferred_qkv_scale_inv: Dict[str, torch.Tensor] = {}
for name, loaded_weight in weights:
is_vision_weight = name.startswith(self._VISION_WEIGHT_PREFIXES)
@@ -1388,7 +1553,11 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
self.config
)
load_mimo_v2_qkv_proj_weight(
name, param, loaded_weight, expected_fused_tp_size
name,
param,
loaded_weight,
expected_fused_tp_size,
deferred_scale_inv=deferred_qkv_scale_inv,
)
continue
@@ -1449,6 +1618,15 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
else:
logger.warning(f"Parameter {name} not found in params_dict")
if deferred_qkv_scale_inv:
expected_fused_tp_size = get_mimo_v2_fused_qkv_expected_tp_size(self.config)
_resolve_deferred_qkv_scale_inv(
params_dict,
deferred_qkv_scale_inv,
expected_fused_tp_size,
config=self.config,
)
def get_embed_and_head(self):
assert (
self.model is not None and self.lm_head is not None
+13
View File
@@ -722,6 +722,18 @@ def get_buffer(name: str, factory: Any) -> Any:
return _CONTEXT.get_buffer(name, factory)
_GLOBAL_DWDP_MANAGER: Any = None
def get_global_dwdp_manager() -> Any:
return _GLOBAL_DWDP_MANAGER
def set_global_dwdp_manager(manager: Any) -> None:
global _GLOBAL_DWDP_MANAGER
_GLOBAL_DWDP_MANAGER = manager
def reset_context() -> None:
"""Clear the context-owned store (unit-test teardown): drop the published
``server_args`` and install fresh ``Flags`` and ``Resources``.
@@ -732,3 +744,4 @@ def reset_context() -> None:
_CONTEXT.flags = Flags()
_CONTEXT.resources = Resources()
_CONTEXT.forward = ForwardFlags()
set_global_dwdp_manager(None)
+68 -1
View File
@@ -953,6 +953,14 @@ class ServerArgs:
aliases=["--decode-context-parallel-size"],
),
] = 1
dwdp_size: A[
int,
Arg(
help="DWDP (Distributed Weight Data Parallelism) group size. "
"When > 1, MoE prefill uses weight prefetch instead of token all-to-all. "
"Must equal tp_size. Only supported with --disaggregation-mode null or prefill.",
),
] = 1
enable_prefill_cp: A[
bool,
"Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.",
@@ -2943,6 +2951,10 @@ class ServerArgs:
# resolution (the declarative registry materializes too late to affect
# it). Inkling opts into full-graph prefill capture here.
self._apply_inkling_prefill_cuda_graph_default()
# must run before _handle_cuda_graph_config and _handle_data_parallelism
self._handle_dwdp()
self._handle_cuda_graph_config()
# Handle device-specific backends.
@@ -4607,7 +4619,7 @@ class ServerArgs:
)
if (
expected_attn_tp_size is not None
and effective_attn_tp_size != expected_attn_tp_size
and expected_attn_tp_size % effective_attn_tp_size != 0
):
raise ValueError(
"MiMoV2ForCausalLM requires effective attention TP "
@@ -5561,6 +5573,61 @@ class ServerArgs:
init_cp_strategy(self)
def _handle_dwdp(self):
if self.dwdp_size <= 1:
return
assert (
self.dwdp_size >= 2
), f"dwdp_size must be >= 2 when enabled, got {self.dwdp_size}"
assert (
self.dwdp_size == self.tp_size
), f"dwdp_size ({self.dwdp_size}) must equal tp_size ({self.tp_size})"
assert self.disaggregation_mode in (
"null",
"prefill",
), "DWDP requires --disaggregation-mode null or prefill"
assert (
not self.enable_eplb
), "EPLB dynamic migration conflicts with static DWDP partitioning"
assert (
self.speculative_algorithm is None
), "DWDP does not support speculative decoding (MTP/draft workers)"
assert self.pp_size == 1, "DWDP requires pp_size == 1"
assert (
not self.enable_two_batch_overlap
), "DWDP's prefetch event protocol does not support two-batch overlap"
if self.disaggregation_mode == "null":
logger.warning(
"DWDP with --disaggregation-mode null: decode steps re-fetch all "
"remote expert weights every step, which is slow. DWDP is "
"recommended only with --disaggregation-mode prefill."
)
self.dp_size = self.dwdp_size
self.enable_dp_attention = True
self.enable_dp_attention_local_control_broadcast = True
self.enable_dp_lm_head = True
self.moe_dense_tp_size = 1
self.ep_size = self.dwdp_size
self.moe_ep_size = self.dwdp_size
self.moe_dp_size = 1
self.moe_a2a_backend = "none"
envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.set(True)
self.disable_cuda_graph = True
logger.info(
f"DWDP enabled: dwdp_size={self.dwdp_size}, "
f"auto-forced dp_size={self.dp_size}, moe_ep_size={self.moe_ep_size}, "
f"moe_dense_tp_size=1, moe_a2a_backend=none, "
f"dp_attention_local_control_broadcast=True, "
f"enable_dp_lm_head=True, SCHEDULER_SKIP_ALL_GATHER=True, "
f"disable_cuda_graph=True"
)
def _handle_data_parallelism(self):
# The dp_size==1 resets moved to the resolution pipeline
# (arg_groups/overrides.py: _data_parallelism_defaults).