Consolidate CUDA VMM allocation helpers (#34199)
This commit is contained in:
+380
-26
@@ -1,7 +1,13 @@
|
|||||||
|
import array
|
||||||
|
import ctypes
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import socket
|
||||||
import struct
|
import struct
|
||||||
|
import tempfile
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
|
from functools import cache
|
||||||
from typing import Any, List, Optional
|
from typing import Any, List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -12,18 +18,43 @@ from sglang.srt.utils import log_info_on_rank0
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_drv = None
|
|
||||||
_FD_HEADER_BYTES = 24
|
_FD_HEADER_BYTES = 24
|
||||||
_FD_SEND_TIMEOUT_S = 120.0
|
_FD_SEND_TIMEOUT_S = 120.0
|
||||||
|
|
||||||
|
try:
|
||||||
|
from cuda.bindings import driver as _drv
|
||||||
|
except ImportError:
|
||||||
|
_drv = None
|
||||||
|
|
||||||
|
if _drv is None:
|
||||||
|
_RECOMMENDED_GRANULARITY = 1
|
||||||
|
else:
|
||||||
|
_RECOMMENDED_GRANULARITY = (
|
||||||
|
_drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
import pynvml
|
||||||
|
except ImportError:
|
||||||
|
pynvml = None
|
||||||
|
|
||||||
|
_NVML_GPU_FABRIC_INFO_V3_TYPE = None
|
||||||
|
_NVML_GPU_FABRIC_INFO_V3_VERSION = None
|
||||||
|
if pynvml is not None:
|
||||||
|
try:
|
||||||
|
_NVML_GPU_FABRIC_INFO_V3_TYPE = pynvml.c_nvmlGpuFabricInfo_v3_t
|
||||||
|
_NVML_GPU_FABRIC_INFO_V3_VERSION = pynvml.nvmlGpuFabricInfo_v3
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# NVML_GPU_FABRIC_STATE_COMPLETED: the GPU has joined its NVLink fabric clique.
|
||||||
|
_NVML_GPU_FABRIC_STATE_COMPLETED = 3
|
||||||
|
|
||||||
|
|
||||||
def _get_cuda_driver():
|
def _get_cuda_driver():
|
||||||
"""Lazily import cuda.bindings.driver (cached after first call)."""
|
"""Return the imported CUDA driver bindings."""
|
||||||
global _drv
|
|
||||||
if _drv is None:
|
if _drv is None:
|
||||||
from cuda.bindings import driver
|
raise ImportError("cuda.bindings.driver is required for CUDA VMM operations")
|
||||||
|
|
||||||
_drv = driver
|
|
||||||
return _drv
|
return _drv
|
||||||
|
|
||||||
|
|
||||||
@@ -38,6 +69,22 @@ def check_drv(result_tuple, label):
|
|||||||
return result_tuple[1] if len(result_tuple) > 1 else None
|
return result_tuple[1] if len(result_tuple) > 1 else None
|
||||||
|
|
||||||
|
|
||||||
|
def tensor_from_pointer(
|
||||||
|
pointer: int,
|
||||||
|
nbytes: int,
|
||||||
|
*,
|
||||||
|
shape=None,
|
||||||
|
dtype: torch.dtype = torch.uint8,
|
||||||
|
device_id: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Use non-owning storage; the caller controls the underlying pages' lifetime."""
|
||||||
|
device = torch.device("cuda", device_id)
|
||||||
|
storage = torch._C._construct_storage_from_data_pointer(pointer, device, nbytes)
|
||||||
|
if shape is None:
|
||||||
|
shape = (nbytes,)
|
||||||
|
return torch.empty(0, dtype=dtype, device=device).set_(storage, 0, shape)
|
||||||
|
|
||||||
|
|
||||||
def is_vmm_pointer(ptr: int) -> bool:
|
def is_vmm_pointer(ptr: int) -> bool:
|
||||||
"""Check if a device pointer is VMM-backed (cuMemCreate/cuMemMap).
|
"""Check if a device pointer is VMM-backed (cuMemCreate/cuMemMap).
|
||||||
|
|
||||||
@@ -112,6 +159,329 @@ def make_rw_access_desc(device_id: int):
|
|||||||
return desc
|
return desc
|
||||||
|
|
||||||
|
|
||||||
|
def _gpu_fabric_clique(device: torch.device):
|
||||||
|
"""Return this GPU's NVLink fabric clique, or ``None`` if not joined."""
|
||||||
|
if pynvml is None:
|
||||||
|
return None
|
||||||
|
cuda_visible_devices = os.environ.get("CUDA_VISIBLE_DEVICES", None)
|
||||||
|
if cuda_visible_devices:
|
||||||
|
device_ids = list(map(int, cuda_visible_devices.split(",")))
|
||||||
|
else:
|
||||||
|
device_ids = list(range(torch.cuda.device_count()))
|
||||||
|
handle = pynvml.nvmlDeviceGetHandleByIndex(device_ids[device.index])
|
||||||
|
if (
|
||||||
|
_NVML_GPU_FABRIC_INFO_V3_TYPE is not None
|
||||||
|
and _NVML_GPU_FABRIC_INFO_V3_VERSION is not None
|
||||||
|
):
|
||||||
|
fabric = _NVML_GPU_FABRIC_INFO_V3_TYPE()
|
||||||
|
fabric.version = _NVML_GPU_FABRIC_INFO_V3_VERSION
|
||||||
|
pynvml.nvmlDeviceGetGpuFabricInfoV(handle, ctypes.byref(fabric))
|
||||||
|
clique_id = fabric.cliqueId
|
||||||
|
else:
|
||||||
|
fabric = pynvml.c_nvmlGpuFabricInfo_t()
|
||||||
|
pynvml.nvmlDeviceGetGpuFabricInfo(handle, ctypes.byref(fabric))
|
||||||
|
clique_id = fabric.partitionId
|
||||||
|
if fabric.state != _NVML_GPU_FABRIC_STATE_COMPLETED:
|
||||||
|
return None
|
||||||
|
return (bytes(fabric.clusterUuid), int(clique_id))
|
||||||
|
|
||||||
|
|
||||||
|
def is_gpu_fabric_ready(device: torch.device) -> bool:
|
||||||
|
"""Whether one CUDA GPU has completed NVLink fabric initialization."""
|
||||||
|
if pynvml is None:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
pynvml.nvmlInit()
|
||||||
|
try:
|
||||||
|
return _gpu_fabric_clique(device) is not None
|
||||||
|
finally:
|
||||||
|
pynvml.nvmlShutdown()
|
||||||
|
except Exception as error:
|
||||||
|
logger.warning("GPU fabric readiness query failed: %r", error)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def allocation_handle_type_name(handle_type: int) -> str:
|
||||||
|
"""Return a stable display name for a CUDA allocation handle type."""
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
fabric = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
|
||||||
|
posix_fd = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
|
||||||
|
if handle_type == fabric:
|
||||||
|
return "FABRIC"
|
||||||
|
if handle_type == posix_fd:
|
||||||
|
return "POSIX_FD"
|
||||||
|
if handle_type == 0:
|
||||||
|
return "NONE"
|
||||||
|
return str(handle_type)
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def get_device_allocation_handle_type(device_id: int) -> int:
|
||||||
|
"""Probe and cache the best supported VMM handle type for one device."""
|
||||||
|
device_id = int(device_id)
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
if not is_gpu_fabric_ready(torch.device("cuda", device_id)):
|
||||||
|
logger.info(
|
||||||
|
"GPU %d has not joined an NVLink fabric clique; probing local "
|
||||||
|
"FABRIC allocation support",
|
||||||
|
device_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
fabric = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
|
||||||
|
posix_fd = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
|
||||||
|
candidates = (fabric, posix_fd, 0)
|
||||||
|
last_error = None
|
||||||
|
for handle_type in candidates:
|
||||||
|
name = allocation_handle_type_name(handle_type)
|
||||||
|
prop = make_device_allocation_prop(
|
||||||
|
device_id,
|
||||||
|
handle_types=handle_type,
|
||||||
|
gpu_direct_rdma=False,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
granularity = get_allocation_granularity(prop)
|
||||||
|
probe_handle = check_drv(
|
||||||
|
drv.cuMemCreate(granularity, prop, 0),
|
||||||
|
f"cuMemCreate({name} probe)",
|
||||||
|
)
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemRelease(probe_handle),
|
||||||
|
f"cuMemRelease({name} probe)",
|
||||||
|
)
|
||||||
|
except RuntimeError as error:
|
||||||
|
last_error = error
|
||||||
|
logger.warning(
|
||||||
|
"CUDA VMM %s backing unavailable on device %d; trying fallback: %s",
|
||||||
|
name,
|
||||||
|
device_id,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
logger.info(
|
||||||
|
"CUDA VMM selected %s backing for device %d",
|
||||||
|
name,
|
||||||
|
device_id,
|
||||||
|
)
|
||||||
|
return handle_type
|
||||||
|
raise RuntimeError("no supported CUDA VMM allocation handle type") from last_error
|
||||||
|
|
||||||
|
|
||||||
|
def make_device_allocation_prop(
|
||||||
|
device_id: int,
|
||||||
|
*,
|
||||||
|
handle_types: int | str | None = "auto",
|
||||||
|
gpu_direct_rdma: bool = False,
|
||||||
|
):
|
||||||
|
"""Build a device allocation prop with automatic or explicit exportability."""
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
if handle_types == "auto":
|
||||||
|
handle_types = get_device_allocation_handle_type(device_id)
|
||||||
|
elif handle_types is None:
|
||||||
|
handle_types = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_NONE
|
||||||
|
elif not isinstance(handle_types, int):
|
||||||
|
raise ValueError("handle_types must be 'auto', an integer, or None")
|
||||||
|
|
||||||
|
handle_types = int(handle_types)
|
||||||
|
valid_handle_types = {
|
||||||
|
int(drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_NONE),
|
||||||
|
int(drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR),
|
||||||
|
int(drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC),
|
||||||
|
}
|
||||||
|
if handle_types not in valid_handle_types:
|
||||||
|
raise ValueError(f"invalid CUDA handle-type value: {handle_types}")
|
||||||
|
|
||||||
|
prop = drv.CUmemAllocationProp()
|
||||||
|
prop.type = drv.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
|
||||||
|
prop.location.type = drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
||||||
|
prop.location.id = int(device_id)
|
||||||
|
prop.requestedHandleTypes = handle_types
|
||||||
|
prop.allocFlags.gpuDirectRDMACapable = int(gpu_direct_rdma)
|
||||||
|
return prop
|
||||||
|
|
||||||
|
|
||||||
|
def get_allocation_granularity(prop, flag=_RECOMMENDED_GRANULARITY) -> int:
|
||||||
|
"""Return allocation granularity for a CUDA policy flag."""
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
return int(
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemGetAllocationGranularity(prop, flag),
|
||||||
|
"cuMemGetAllocationGranularity",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def get_device_granularity(device_id: int) -> int:
|
||||||
|
"""Granularity for this device's default allocations. Cached: it is a device
|
||||||
|
constant, and callers that size a reservation must agree with the one that
|
||||||
|
maps into it."""
|
||||||
|
device_id = int(device_id)
|
||||||
|
return get_allocation_granularity(make_device_allocation_prop(device_id))
|
||||||
|
|
||||||
|
|
||||||
|
def align_up(value: int, alignment: int) -> int:
|
||||||
|
"""Round ``value`` up to a positive byte ``alignment``."""
|
||||||
|
return (int(value) + alignment - 1) // alignment * alignment
|
||||||
|
|
||||||
|
|
||||||
|
def align_down(value: int, alignment: int) -> int:
|
||||||
|
"""Round ``value`` down to a positive byte ``alignment``."""
|
||||||
|
return int(value) // alignment * alignment
|
||||||
|
|
||||||
|
|
||||||
|
class VmmReservation:
|
||||||
|
"""Own a VA reservation, its mappings, and their teardown order."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
size: int,
|
||||||
|
prop,
|
||||||
|
device_id: int,
|
||||||
|
*,
|
||||||
|
alignment: int = 0,
|
||||||
|
requested_address: int = 0,
|
||||||
|
) -> None:
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
self.size = int(size)
|
||||||
|
self._prop = prop
|
||||||
|
self._access_descs = [make_rw_access_desc(int(device_id))]
|
||||||
|
self.base = int(
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemAddressReserve(
|
||||||
|
self.size,
|
||||||
|
int(alignment),
|
||||||
|
int(requested_address),
|
||||||
|
0,
|
||||||
|
),
|
||||||
|
"cuMemAddressReserve(local)",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self._mappings = []
|
||||||
|
self._closed = False
|
||||||
|
|
||||||
|
def map(
|
||||||
|
self,
|
||||||
|
offset: int,
|
||||||
|
size: int,
|
||||||
|
*,
|
||||||
|
retain_handle: bool,
|
||||||
|
):
|
||||||
|
"""Create and map local memory at ``base + offset``."""
|
||||||
|
if self._closed:
|
||||||
|
raise RuntimeError("VmmReservation.map after close")
|
||||||
|
offset, size = int(offset), int(size)
|
||||||
|
if offset < 0 or size <= 0 or offset + size > self.size:
|
||||||
|
raise ValueError(
|
||||||
|
f"mapping [{offset}, {offset + size}) is outside reservation "
|
||||||
|
f"[0, {self.size})"
|
||||||
|
)
|
||||||
|
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
address = self.base + offset
|
||||||
|
handle = check_drv(drv.cuMemCreate(size, self._prop, 0), "cuMemCreate(local)")
|
||||||
|
mapped = False
|
||||||
|
try:
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemMap(address, size, 0, handle, 0),
|
||||||
|
"cuMemMap(local)",
|
||||||
|
)
|
||||||
|
mapped = True
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemSetAccess(
|
||||||
|
address,
|
||||||
|
size,
|
||||||
|
self._access_descs,
|
||||||
|
len(self._access_descs),
|
||||||
|
),
|
||||||
|
"cuMemSetAccess(local)",
|
||||||
|
)
|
||||||
|
if not retain_handle:
|
||||||
|
check_drv(drv.cuMemRelease(handle), "cuMemRelease(local)")
|
||||||
|
handle = None
|
||||||
|
except BaseException as error:
|
||||||
|
cleanup_errors = []
|
||||||
|
if mapped:
|
||||||
|
try:
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemUnmap(address, size), "cuMemUnmap(local rollback)"
|
||||||
|
)
|
||||||
|
except BaseException as cleanup_error:
|
||||||
|
cleanup_errors.append(cleanup_error)
|
||||||
|
if handle is not None:
|
||||||
|
try:
|
||||||
|
check_drv(drv.cuMemRelease(handle), "cuMemRelease(local rollback)")
|
||||||
|
except BaseException as cleanup_error:
|
||||||
|
cleanup_errors.append(cleanup_error)
|
||||||
|
if cleanup_errors:
|
||||||
|
error.add_note(
|
||||||
|
f"{len(cleanup_errors)} CUDA VMM rollback operation(s) also failed"
|
||||||
|
)
|
||||||
|
raise error from cleanup_errors[0]
|
||||||
|
raise
|
||||||
|
|
||||||
|
self._mappings.append((address, size, handle))
|
||||||
|
return handle
|
||||||
|
|
||||||
|
def map_existing(self, offset: int, size: int, handle) -> None:
|
||||||
|
"""Map a caller-owned physical allocation into this reservation."""
|
||||||
|
if self._closed:
|
||||||
|
raise RuntimeError("VmmReservation.map_existing after close")
|
||||||
|
offset, size = int(offset), int(size)
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
address = self.base + offset
|
||||||
|
mapped = False
|
||||||
|
try:
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemMap(address, size, 0, handle, 0),
|
||||||
|
"cuMemMap(existing)",
|
||||||
|
)
|
||||||
|
mapped = True
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemSetAccess(
|
||||||
|
address,
|
||||||
|
size,
|
||||||
|
self._access_descs,
|
||||||
|
len(self._access_descs),
|
||||||
|
),
|
||||||
|
"cuMemSetAccess(existing)",
|
||||||
|
)
|
||||||
|
except BaseException as error:
|
||||||
|
if mapped:
|
||||||
|
try:
|
||||||
|
check_drv(
|
||||||
|
drv.cuMemUnmap(address, size),
|
||||||
|
"cuMemUnmap(existing rollback)",
|
||||||
|
)
|
||||||
|
except BaseException as cleanup_error:
|
||||||
|
error.add_note("CUDA VMM alias rollback also failed")
|
||||||
|
raise error from cleanup_error
|
||||||
|
raise
|
||||||
|
|
||||||
|
self._mappings.append((address, size, None))
|
||||||
|
|
||||||
|
def close(self, *, release_handles: bool = True) -> None:
|
||||||
|
"""Unmap allocations, optionally release retained handles, and free VA."""
|
||||||
|
if self._closed:
|
||||||
|
return
|
||||||
|
self._closed = True
|
||||||
|
drv = _get_cuda_driver()
|
||||||
|
while self._mappings:
|
||||||
|
address, size, handle = self._mappings.pop()
|
||||||
|
err = drv.cuMemUnmap(address, size)
|
||||||
|
err = err[0] if isinstance(err, tuple) else err
|
||||||
|
if err != drv.CUresult.CUDA_SUCCESS:
|
||||||
|
logger.warning("cuMemUnmap(local) -> %s", err)
|
||||||
|
if release_handles and handle is not None:
|
||||||
|
err = drv.cuMemRelease(handle)
|
||||||
|
err = err[0] if isinstance(err, tuple) else err
|
||||||
|
if err != drv.CUresult.CUDA_SUCCESS:
|
||||||
|
logger.warning("cuMemRelease(local) -> %s", err)
|
||||||
|
err = drv.cuMemAddressFree(self.base, self.size)
|
||||||
|
err = err[0] if isinstance(err, tuple) else err
|
||||||
|
if err != drv.CUresult.CUDA_SUCCESS:
|
||||||
|
logger.warning("cuMemAddressFree(local) -> %s", err)
|
||||||
|
|
||||||
|
|
||||||
def all_ranks_ok(group: ProcessGroup, ok: bool) -> bool:
|
def all_ranks_ok(group: ProcessGroup, ok: bool) -> bool:
|
||||||
"""True iff ``ok`` holds on every rank in ``group`` (BAND all-reduce)."""
|
"""True iff ``ok`` holds on every rank in ``group`` (BAND all-reduce)."""
|
||||||
flag = torch.tensor([1 if ok else 0], dtype=torch.int32)
|
flag = torch.tensor([1 if ok else 0], dtype=torch.int32)
|
||||||
@@ -133,9 +503,6 @@ def release_mappings(mappings) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None:
|
def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None:
|
||||||
import array
|
|
||||||
import socket
|
|
||||||
|
|
||||||
fds = array.array("i", [int(fd)])
|
fds = array.array("i", [int(fd)])
|
||||||
header = struct.pack("<QQQ", int(src_rank), int(base_idx), 1)
|
header = struct.pack("<QQQ", int(src_rank), int(base_idx), 1)
|
||||||
sent = sock.sendmsg(
|
sent = sock.sendmsg(
|
||||||
@@ -147,9 +514,6 @@ def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def _recv_fd(sock):
|
def _recv_fd(sock):
|
||||||
import array
|
|
||||||
import socket
|
|
||||||
|
|
||||||
fd_item_size = array.array("i").itemsize
|
fd_item_size = array.array("i").itemsize
|
||||||
data, ancdata, _, _ = sock.recvmsg(
|
data, ancdata, _, _ = sock.recvmsg(
|
||||||
_FD_HEADER_BYTES, socket.CMSG_SPACE(fd_item_size)
|
_FD_HEADER_BYTES, socket.CMSG_SPACE(fd_item_size)
|
||||||
@@ -253,11 +617,7 @@ def exchange_posix_fds(
|
|||||||
socket. Returns ``{(src_rank, base_idx): fd}`` for every peer. The caller
|
socket. Returns ``{(src_rank, base_idx): fd}`` for every peer. The caller
|
||||||
owns the received fds and must close them.
|
owns the received fds and must close them.
|
||||||
"""
|
"""
|
||||||
import socket
|
sock_kind = socket.SOCK_SEQPACKET
|
||||||
import tempfile
|
|
||||||
import threading
|
|
||||||
|
|
||||||
sock_kind = getattr(socket, "SOCK_SEQPACKET", socket.SOCK_STREAM)
|
|
||||||
sock_dir = tempfile.mkdtemp(prefix="sgl_ar_fd_")
|
sock_dir = tempfile.mkdtemp(prefix="sgl_ar_fd_")
|
||||||
sock_path = os.path.join(sock_dir, f"rank_{rank}.sock")
|
sock_path = os.path.join(sock_dir, f"rank_{rank}.sock")
|
||||||
server = socket.socket(socket.AF_UNIX, sock_kind)
|
server = socket.socket(socket.AF_UNIX, sock_kind)
|
||||||
@@ -382,13 +742,7 @@ def import_and_map_alloc(
|
|||||||
drv.cuMemGetAllocationPropertiesFromHandle(imp_h),
|
drv.cuMemGetAllocationPropertiesFromHandle(imp_h),
|
||||||
"cuMemGetAllocationPropertiesFromHandle",
|
"cuMemGetAllocationPropertiesFromHandle",
|
||||||
)
|
)
|
||||||
gran = check_drv(
|
gran = get_allocation_granularity(prop)
|
||||||
drv.cuMemGetAllocationGranularity(
|
|
||||||
prop,
|
|
||||||
drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED,
|
|
||||||
),
|
|
||||||
"cuMemGetAllocationGranularity",
|
|
||||||
)
|
|
||||||
va = check_drv(
|
va = check_drv(
|
||||||
drv.cuMemAddressReserve(alloc_size, int(gran), 0, 0), "cuMemAddressReserve"
|
drv.cuMemAddressReserve(alloc_size, int(gran), 0, 0), "cuMemAddressReserve"
|
||||||
)
|
)
|
||||||
@@ -447,8 +801,8 @@ class VmmGraphInputManager:
|
|||||||
VMM-compatible path for expandable_segments. The C++ side deduplicates
|
VMM-compatible path for expandable_segments. The C++ side deduplicates
|
||||||
graph capture pointers into unique base allocations via cuMemGetAddressRange.
|
graph capture pointers into unique base allocations via cuMemGetAddressRange.
|
||||||
Python exports handles for each unique base, imports + cuMemMaps peer
|
Python exports handles for each unique base, imports + cuMemMaps peer
|
||||||
allocations, then registers the peer VAs. FABRIC handles are preferred;
|
allocations, then registers the peer virtual addresses. FABRIC handles are
|
||||||
POSIX file descriptors are used when FABRIC is unavailable.
|
preferred; POSIX file descriptors are used when FABRIC is unavailable.
|
||||||
"""
|
"""
|
||||||
FABRIC_HANDLE_BYTES = 64
|
FABRIC_HANDLE_BYTES = 64
|
||||||
MAX_VMM_BASES = 4096
|
MAX_VMM_BASES = 4096
|
||||||
@@ -19,6 +19,7 @@ import torch.distributed as dist
|
|||||||
import torch.multiprocessing as mp
|
import torch.multiprocessing as mp
|
||||||
from typing_extensions import ParamSpec
|
from typing_extensions import ParamSpec
|
||||||
|
|
||||||
|
from sglang.srt.cuda_vmm_utils import _gpu_fabric_clique
|
||||||
from sglang.srt.distributed.device_communicators.cuda_wrapper import CudaRTLibrary
|
from sglang.srt.distributed.device_communicators.cuda_wrapper import CudaRTLibrary
|
||||||
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
||||||
from sglang.srt.environ import envs as sglang_envs
|
from sglang.srt.environ import envs as sglang_envs
|
||||||
@@ -391,27 +392,6 @@ def is_full_nvlink(physical_device_ids: List[int], world_size: int) -> bool:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
# NVML_GPU_FABRIC_STATE_COMPLETED: the GPU has joined its NVLink fabric clique.
|
|
||||||
_NVML_GPU_FABRIC_STATE_COMPLETED = 3
|
|
||||||
|
|
||||||
|
|
||||||
def _gpu_fabric_clique(device: torch.device):
|
|
||||||
"""(cluster_uuid, clique_id) of the local GPU's NVLink fabric clique, or None if
|
|
||||||
the GPU has not joined a fabric (single-node box / fabric init incomplete)."""
|
|
||||||
cuda_visible_devices = os.environ.get("CUDA_VISIBLE_DEVICES", None)
|
|
||||||
if cuda_visible_devices:
|
|
||||||
device_ids = list(map(int, cuda_visible_devices.split(",")))
|
|
||||||
else:
|
|
||||||
device_ids = list(range(torch.cuda.device_count()))
|
|
||||||
handle = pynvml.nvmlDeviceGetHandleByIndex(device_ids[device.index])
|
|
||||||
fabric = pynvml.c_nvmlGpuFabricInfo_v3_t()
|
|
||||||
fabric.version = pynvml.nvmlGpuFabricInfo_v3
|
|
||||||
pynvml.nvmlDeviceGetGpuFabricInfoV(handle, ctypes.byref(fabric))
|
|
||||||
if fabric.state != _NVML_GPU_FABRIC_STATE_COMPLETED:
|
|
||||||
return None
|
|
||||||
return (bytes(fabric.clusterUuid), int(fabric.cliqueId))
|
|
||||||
|
|
||||||
|
|
||||||
@with_nvml_context
|
@with_nvml_context
|
||||||
def is_one_nvlink_clique(
|
def is_one_nvlink_clique(
|
||||||
group: torch.distributed.ProcessGroup, device: torch.device
|
group: torch.distributed.ProcessGroup, device: torch.device
|
||||||
|
|||||||
@@ -31,6 +31,11 @@ from sglang.kernels.ops.communication.all_reduce import (
|
|||||||
IPCManager,
|
IPCManager,
|
||||||
custom_all_reduce,
|
custom_all_reduce,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
|
VmmGraphInputManager,
|
||||||
|
compute_graph_capture_bases,
|
||||||
|
is_vmm_pointer,
|
||||||
|
)
|
||||||
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
from sglang.srt.distributed.parallel_state import in_the_same_node_as
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
@@ -43,11 +48,6 @@ from .custom_all_reduce_utils import (
|
|||||||
is_one_nvlink_clique,
|
is_one_nvlink_clique,
|
||||||
is_weak_contiguous,
|
is_weak_contiguous,
|
||||||
)
|
)
|
||||||
from .vmm_utils import (
|
|
||||||
VmmGraphInputManager,
|
|
||||||
compute_graph_capture_bases,
|
|
||||||
is_vmm_pointer,
|
|
||||||
)
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from typing import Dict, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.layers.moe.dwdp.vmm import align_down, align_up
|
from sglang.srt.cuda_vmm_utils import align_down, align_up
|
||||||
|
|
||||||
# one (start, end_capped) expert range per peer DWDP rank
|
# one (start, end_capped) expert range per peer DWDP rank
|
||||||
PeerRanges = List[Tuple[int, int]]
|
PeerRanges = List[Tuple[int, int]]
|
||||||
|
|||||||
@@ -4,14 +4,16 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Dict, List, Optional, Tuple
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
from sglang.srt.layers.moe.dwdp.vmm import (
|
from cuda.bindings import driver as cuda
|
||||||
|
|
||||||
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
|
VmmReservation,
|
||||||
align_up,
|
align_up,
|
||||||
create_local_handle,
|
check_drv,
|
||||||
get_allocation_granularity,
|
get_device_granularity,
|
||||||
map_handle,
|
make_device_allocation_prop,
|
||||||
release_handle,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -31,7 +33,8 @@ class PagePool:
|
|||||||
page_size: Optional[int] = None,
|
page_size: Optional[int] = None,
|
||||||
):
|
):
|
||||||
self._device_id = device_id
|
self._device_id = device_id
|
||||||
self._granularity = granularity or get_allocation_granularity(device_id)
|
self._granularity = granularity or get_device_granularity(device_id)
|
||||||
|
self._prop = make_device_allocation_prop(device_id, handle_types=None)
|
||||||
|
|
||||||
if page_size is None:
|
if page_size is None:
|
||||||
self._page_size = self.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
|
self._page_size = self.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
|
||||||
@@ -49,8 +52,15 @@ class PagePool:
|
|||||||
for slot_idx, num_pages in enumerate(self._slot_pages):
|
for slot_idx, num_pages in enumerate(self._slot_pages):
|
||||||
handles = []
|
handles = []
|
||||||
for _ in range(num_pages):
|
for _ in range(num_pages):
|
||||||
h = create_local_handle(self._page_size, device_id)
|
reservation = VmmReservation(
|
||||||
handles.append(h)
|
self._page_size,
|
||||||
|
self._prop,
|
||||||
|
device_id,
|
||||||
|
alignment=self._granularity,
|
||||||
|
)
|
||||||
|
handle = reservation.map(0, self._page_size, retain_handle=True)
|
||||||
|
reservation.close(release_handles=False)
|
||||||
|
handles.append(int(handle))
|
||||||
self._page_handles.append(handles)
|
self._page_handles.append(handles)
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"PagePool slot {slot_idx}: {num_pages} pages × {self._page_size} B"
|
f"PagePool slot {slot_idx}: {num_pages} pages × {self._page_size} B"
|
||||||
@@ -78,21 +88,21 @@ class PagePool:
|
|||||||
def map_pages(
|
def map_pages(
|
||||||
self,
|
self,
|
||||||
slot: int,
|
slot: int,
|
||||||
va_start: int,
|
reservation: VmmReservation,
|
||||||
|
offset: int,
|
||||||
size: int,
|
size: int,
|
||||||
page_offset: int = 0,
|
page_offset: int = 0,
|
||||||
) -> List[Tuple[int, int]]:
|
) -> None:
|
||||||
# does NOT call set_access; caller must set access on the whole composite VA
|
|
||||||
aligned_size = align_up(size, self._page_size)
|
aligned_size = align_up(size, self._page_size)
|
||||||
num_pages_needed = aligned_size // self._page_size
|
num_pages_needed = aligned_size // self._page_size
|
||||||
|
|
||||||
mappings = []
|
|
||||||
for i in range(num_pages_needed):
|
for i in range(num_pages_needed):
|
||||||
va = va_start + i * self._page_size
|
|
||||||
handle = self._page_handles[slot][page_offset + i]
|
handle = self._page_handles[slot][page_offset + i]
|
||||||
map_handle(va, self._page_size, handle, offset=0)
|
reservation.map_existing(
|
||||||
mappings.append((va, self._page_size))
|
offset + i * self._page_size,
|
||||||
return mappings
|
self._page_size,
|
||||||
|
handle,
|
||||||
|
)
|
||||||
|
|
||||||
def release(self) -> None:
|
def release(self) -> None:
|
||||||
if self._released:
|
if self._released:
|
||||||
@@ -100,7 +110,7 @@ class PagePool:
|
|||||||
self._released = True
|
self._released = True
|
||||||
for handles in self._page_handles:
|
for handles in self._page_handles:
|
||||||
for h in handles:
|
for h in handles:
|
||||||
release_handle(h)
|
check_drv(cuda.cuMemRelease(h), "cuMemRelease")
|
||||||
self._page_handles = [[], []]
|
self._page_handles = [[], []]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -11,30 +11,23 @@ import torch
|
|||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
from cuda.bindings import driver as cuda
|
from cuda.bindings import driver as cuda
|
||||||
|
|
||||||
from sglang.srt.distributed.device_communicators.vmm_utils import (
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
|
VmmReservation,
|
||||||
|
align_down,
|
||||||
|
align_up,
|
||||||
check_drv,
|
check_drv,
|
||||||
exchange_posix_fds,
|
exchange_posix_fds,
|
||||||
export_shareable_handles,
|
export_shareable_handles,
|
||||||
|
get_device_granularity,
|
||||||
import_peer_handle,
|
import_peer_handle,
|
||||||
|
make_device_allocation_prop,
|
||||||
|
tensor_from_pointer,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.dwdp.layout import (
|
from sglang.srt.layers.moe.dwdp.layout import (
|
||||||
DwdpExpertLayout,
|
DwdpExpertLayout,
|
||||||
LayerWeightSpecs,
|
LayerWeightSpecs,
|
||||||
MnnvlHandleSet,
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -54,7 +47,8 @@ def _copy_local_weights_to_handles(
|
|||||||
layout: DwdpExpertLayout,
|
layout: DwdpExpertLayout,
|
||||||
device_id: int,
|
device_id: int,
|
||||||
) -> Tuple[Dict[Tuple[int, str], int], Dict[Tuple[int, str], int]]:
|
) -> Tuple[Dict[Tuple[int, str], int], Dict[Tuple[int, str], int]]:
|
||||||
granularity = get_allocation_granularity(device_id)
|
granularity = get_device_granularity(device_id)
|
||||||
|
prop = make_device_allocation_prop(device_id)
|
||||||
handles: Dict[Tuple[int, str], int] = {}
|
handles: Dict[Tuple[int, str], int] = {}
|
||||||
sizes: Dict[Tuple[int, str], int] = {}
|
sizes: Dict[Tuple[int, str], int] = {}
|
||||||
|
|
||||||
@@ -69,21 +63,17 @@ def _copy_local_weights_to_handles(
|
|||||||
phys_size = page_end - page_start
|
phys_size = page_end - page_start
|
||||||
data_offset = local_start_bytes - page_start
|
data_offset = local_start_bytes - page_start
|
||||||
|
|
||||||
handle = create_fabric_handle(phys_size, device_id)
|
reservation = VmmReservation(phys_size, prop, device_id, alignment=granularity)
|
||||||
|
handle = int(reservation.map(0, phys_size, retain_handle=True))
|
||||||
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()
|
nbytes = param.numel() * param.element_size()
|
||||||
check_drv(
|
check_drv(
|
||||||
cuda.cuMemcpyDtoD(temp_va + data_offset, param.data_ptr(), nbytes),
|
cuda.cuMemcpyDtoD(reservation.base + data_offset, param.data_ptr(), nbytes),
|
||||||
"cuMemcpyDtoD",
|
"cuMemcpyDtoD",
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
unmap_va(temp_va, phys_size)
|
reservation.close(release_handles=False)
|
||||||
free_va(temp_va, phys_size)
|
|
||||||
|
|
||||||
param.untyped_storage().resize_(0)
|
param.untyped_storage().resize_(0)
|
||||||
|
|
||||||
@@ -105,7 +95,7 @@ class DWDPTransport:
|
|||||||
self._handle_set: Optional[MnnvlHandleSet] = None
|
self._handle_set: Optional[MnnvlHandleSet] = None
|
||||||
self._peer_views: Dict[Tuple[int, int, str], torch.Tensor] = {}
|
self._peer_views: Dict[Tuple[int, int, str], torch.Tensor] = {}
|
||||||
self._imported_handles: List[int] = []
|
self._imported_handles: List[int] = []
|
||||||
self._peer_va_regions: List[Tuple[int, int]] = []
|
self._peer_reservations: List[VmmReservation] = []
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
@@ -144,7 +134,8 @@ class DWDPTransport:
|
|||||||
device_id: int,
|
device_id: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
cpu_group = group.cpu_group
|
cpu_group = group.cpu_group
|
||||||
granularity = get_allocation_granularity(device_id)
|
granularity = get_device_granularity(device_id)
|
||||||
|
prop = make_device_allocation_prop(device_id)
|
||||||
|
|
||||||
handle_list = [self._handle_set.get_handle(li, n) for li, n in sorted_keys]
|
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(
|
fabric_handles, local_posix_fds, use_fabric = export_shareable_handles(
|
||||||
@@ -199,14 +190,19 @@ class DWDPTransport:
|
|||||||
peer_phys_size = peer_page_end - peer_page_start
|
peer_phys_size = peer_page_end - peer_page_start
|
||||||
peer_data_offset = peer_start_bytes - peer_page_start
|
peer_data_offset = peer_start_bytes - peer_page_start
|
||||||
|
|
||||||
peer_va = reserve_va(peer_phys_size, granularity)
|
peer_reservation = VmmReservation(
|
||||||
map_handle(peer_va, peer_phys_size, int(peer_handle))
|
peer_phys_size,
|
||||||
set_access(peer_va, peer_phys_size, device_id)
|
prop,
|
||||||
self._peer_va_regions.append((peer_va, peer_phys_size))
|
device_id,
|
||||||
|
alignment=granularity,
|
||||||
|
)
|
||||||
|
peer_reservation.map_existing(0, peer_phys_size, int(peer_handle))
|
||||||
|
self._peer_reservations.append(peer_reservation)
|
||||||
|
|
||||||
num_peer_experts = peer_end - peer_start
|
num_peer_experts = peer_end - peer_start
|
||||||
peer_tensor = tensor_from_ptr(
|
peer_tensor = tensor_from_pointer(
|
||||||
ptr=peer_va + peer_data_offset,
|
peer_reservation.base + peer_data_offset,
|
||||||
|
peer_end_bytes - peer_start_bytes,
|
||||||
shape=(num_peer_experts,) + spec.full_shape[1:],
|
shape=(num_peer_experts,) + spec.full_shape[1:],
|
||||||
dtype=spec.dtype,
|
dtype=spec.dtype,
|
||||||
device_id=device_id,
|
device_id=device_id,
|
||||||
@@ -226,13 +222,12 @@ class DWDPTransport:
|
|||||||
return self._peer_views
|
return self._peer_views
|
||||||
|
|
||||||
def release(self) -> None:
|
def release(self) -> None:
|
||||||
for va, size in self._peer_va_regions:
|
for reservation in self._peer_reservations:
|
||||||
unmap_va(va, size)
|
reservation.close()
|
||||||
free_va(va, size)
|
self._peer_reservations.clear()
|
||||||
self._peer_va_regions.clear()
|
|
||||||
|
|
||||||
for h in self._imported_handles:
|
for h in self._imported_handles:
|
||||||
release_handle(h)
|
check_drv(cuda.cuMemRelease(h), "cuMemRelease")
|
||||||
self._imported_handles.clear()
|
self._imported_handles.clear()
|
||||||
|
|
||||||
self._peer_views.clear()
|
self._peer_views.clear()
|
||||||
|
|||||||
@@ -1,258 +0,0 @@
|
|||||||
"""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
|
|
||||||
@@ -8,6 +8,12 @@ from typing import Dict, List, Optional, Tuple
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
|
VmmReservation,
|
||||||
|
get_device_granularity,
|
||||||
|
make_device_allocation_prop,
|
||||||
|
tensor_from_pointer,
|
||||||
|
)
|
||||||
from sglang.srt.layers.moe.dwdp.layout import (
|
from sglang.srt.layers.moe.dwdp.layout import (
|
||||||
EdgeInfo,
|
EdgeInfo,
|
||||||
LayerWeightSpecs,
|
LayerWeightSpecs,
|
||||||
@@ -15,15 +21,6 @@ from sglang.srt.layers.moe.dwdp.layout import (
|
|||||||
PageAlignedLayout,
|
PageAlignedLayout,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.dwdp.page_pool import PagePool, compute_slot_sizes
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -44,7 +41,8 @@ class WeightBuffer:
|
|||||||
self._local_end = local_end
|
self._local_end = local_end
|
||||||
self._dwdp_size = dwdp_size
|
self._dwdp_size = dwdp_size
|
||||||
self._device_id = device_id
|
self._device_id = device_id
|
||||||
self._granularity = get_allocation_granularity(device_id)
|
self._granularity = get_device_granularity(device_id)
|
||||||
|
self._prop = make_device_allocation_prop(device_id, handle_types=None)
|
||||||
self._pool_page_size = PagePool.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
|
self._pool_page_size = PagePool.DEFAULT_PAGE_SIZE_MULTIPLIER * self._granularity
|
||||||
self._page_pool: Optional[PagePool] = None
|
self._page_pool: Optional[PagePool] = None
|
||||||
self._moe_layer_indices = sorted(layer_weight_specs.keys())
|
self._moe_layer_indices = sorted(layer_weight_specs.keys())
|
||||||
@@ -53,8 +51,7 @@ class WeightBuffer:
|
|||||||
self._remote_slices: Dict[
|
self._remote_slices: Dict[
|
||||||
int, Dict[str, List[Tuple[torch.Tensor, int, int]]]
|
int, Dict[str, List[Tuple[torch.Tensor, int, int]]]
|
||||||
] = {}
|
] = {}
|
||||||
self._mappings: Dict[int, List[Tuple[int, int]]] = {}
|
self._reservations: Dict[int, List[VmmReservation]] = {}
|
||||||
self._va_regions: Dict[int, List[Tuple[int, int]]] = {}
|
|
||||||
self._released = False
|
self._released = False
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -105,8 +102,7 @@ class WeightBuffer:
|
|||||||
|
|
||||||
self._tensors[layer_idx] = {}
|
self._tensors[layer_idx] = {}
|
||||||
self._remote_slices[layer_idx] = {}
|
self._remote_slices[layer_idx] = {}
|
||||||
self._mappings[layer_idx] = []
|
self._reservations[layer_idx] = []
|
||||||
self._va_regions[layer_idx] = []
|
|
||||||
|
|
||||||
page_pool_offset = 0
|
page_pool_offset = 0
|
||||||
|
|
||||||
@@ -114,40 +110,41 @@ class WeightBuffer:
|
|||||||
spec = weight_specs[name]
|
spec = weight_specs[name]
|
||||||
handle = self._handles.get_handle(layer_idx, name)
|
handle = self._handles.get_handle(layer_idx, name)
|
||||||
|
|
||||||
va_base = reserve_va(layout.total_size, self._granularity)
|
reservation = VmmReservation(
|
||||||
self._va_regions[layer_idx].append((va_base, layout.total_size))
|
layout.total_size,
|
||||||
all_maps = self._mappings[layer_idx]
|
self._prop,
|
||||||
|
self._device_id,
|
||||||
|
alignment=self._granularity,
|
||||||
|
)
|
||||||
|
self._reservations[layer_idx].append(reservation)
|
||||||
|
va_base = reservation.base
|
||||||
|
|
||||||
if layout.pre_size > 0:
|
if layout.pre_size > 0:
|
||||||
pre_maps = self._page_pool.map_pages(
|
self._page_pool.map_pages(
|
||||||
slot=buf_slot,
|
slot=buf_slot,
|
||||||
va_start=va_base,
|
reservation=reservation,
|
||||||
|
offset=0,
|
||||||
size=layout.pre_size,
|
size=layout.pre_size,
|
||||||
page_offset=page_pool_offset,
|
page_offset=page_pool_offset,
|
||||||
)
|
)
|
||||||
all_maps.extend(pre_maps)
|
|
||||||
page_pool_offset += layout.pre_pages
|
page_pool_offset += layout.pre_pages
|
||||||
|
|
||||||
mnnvl_va = va_base + layout.pre_size
|
reservation.map_existing(layout.pre_size, layout.mnnvl_size, handle)
|
||||||
map_handle(mnnvl_va, layout.mnnvl_size, handle, offset=0)
|
|
||||||
all_maps.append((mnnvl_va, layout.mnnvl_size))
|
|
||||||
|
|
||||||
if layout.post_size > 0:
|
if layout.post_size > 0:
|
||||||
post_va = mnnvl_va + layout.mnnvl_size
|
self._page_pool.map_pages(
|
||||||
post_maps = self._page_pool.map_pages(
|
|
||||||
slot=buf_slot,
|
slot=buf_slot,
|
||||||
va_start=post_va,
|
reservation=reservation,
|
||||||
|
offset=layout.pre_size + layout.mnnvl_size,
|
||||||
size=layout.post_size,
|
size=layout.post_size,
|
||||||
page_offset=page_pool_offset,
|
page_offset=page_pool_offset,
|
||||||
)
|
)
|
||||||
all_maps.extend(post_maps)
|
|
||||||
page_pool_offset += layout.post_pages
|
page_pool_offset += layout.post_pages
|
||||||
|
|
||||||
set_access(va_base, layout.total_size, self._device_id)
|
|
||||||
|
|
||||||
tensor_start = va_base + layout.pre_padding
|
tensor_start = va_base + layout.pre_padding
|
||||||
full_tensor = tensor_from_ptr(
|
full_tensor = tensor_from_pointer(
|
||||||
ptr=tensor_start,
|
tensor_start,
|
||||||
|
layout.num_experts * layout.expert_bytes,
|
||||||
shape=spec.full_shape,
|
shape=spec.full_shape,
|
||||||
dtype=spec.dtype,
|
dtype=spec.dtype,
|
||||||
device_id=self._device_id,
|
device_id=self._device_id,
|
||||||
@@ -206,14 +203,10 @@ class WeightBuffer:
|
|||||||
if self._released:
|
if self._released:
|
||||||
return
|
return
|
||||||
self._released = True
|
self._released = True
|
||||||
for li, maps in self._mappings.items():
|
for reservations in self._reservations.values():
|
||||||
for va, sz in maps:
|
for reservation in reservations:
|
||||||
unmap_va(va, sz)
|
reservation.close()
|
||||||
for li, regions in self._va_regions.items():
|
self._reservations.clear()
|
||||||
for va, sz in regions:
|
|
||||||
free_va(va, sz)
|
|
||||||
self._mappings.clear()
|
|
||||||
self._va_regions.clear()
|
|
||||||
self._tensors.clear()
|
self._tensors.clear()
|
||||||
self._remote_slices.clear()
|
self._remote_slices.clear()
|
||||||
if self._page_pool is not None:
|
if self._page_pool is not None:
|
||||||
|
|||||||
@@ -11,52 +11,19 @@ import torch
|
|||||||
import torch.utils.cpp_extension
|
import torch.utils.cpp_extension
|
||||||
from torch.cuda.memory import CUDAPluggableAllocator
|
from torch.cuda.memory import CUDAPluggableAllocator
|
||||||
|
|
||||||
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
|
VmmReservation,
|
||||||
|
align_up,
|
||||||
|
allocation_handle_type_name,
|
||||||
|
get_device_granularity,
|
||||||
|
make_device_allocation_prop,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.memory_pool import KvBufferDesc
|
from sglang.srt.mem_cache.memory_pool import KvBufferDesc
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_drv = None
|
|
||||||
|
|
||||||
|
|
||||||
def _driver():
|
|
||||||
global _drv
|
|
||||||
if _drv is None:
|
|
||||||
from cuda.bindings import driver
|
|
||||||
|
|
||||||
_drv = driver
|
|
||||||
return _drv
|
|
||||||
|
|
||||||
|
|
||||||
def _check(result, label: str):
|
|
||||||
drv = _driver()
|
|
||||||
err = result[0] if isinstance(result, tuple) else result
|
|
||||||
if err != drv.CUresult.CUDA_SUCCESS:
|
|
||||||
raise RuntimeError(f"{label} failed: {err}")
|
|
||||||
return result[1] if isinstance(result, tuple) and len(result) > 1 else None
|
|
||||||
|
|
||||||
|
|
||||||
def align_up(value: int, alignment: int) -> int:
|
|
||||||
return (value + alignment - 1) // alignment * alignment
|
|
||||||
|
|
||||||
|
|
||||||
def query_granularity(device_id: int) -> int:
|
|
||||||
"""Minimum CUDA virtual-memory allocation granularity (bytes) for ``device_id``."""
|
|
||||||
drv = _driver()
|
|
||||||
prop = drv.CUmemAllocationProp()
|
|
||||||
prop.type = drv.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
|
|
||||||
prop.location.type = drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
|
||||||
prop.location.id = int(device_id)
|
|
||||||
return int(
|
|
||||||
_check(
|
|
||||||
drv.cuMemGetAllocationGranularity(
|
|
||||||
prop,
|
|
||||||
drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_MINIMUM,
|
|
||||||
),
|
|
||||||
"cuMemGetAllocationGranularity",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Bump allocator: hands back base+cursor, bounded by the RESERVED size (not the
|
# Bump allocator: hands back base+cursor, bounded by the RESERVED size (not the
|
||||||
# committed watermark) so upper-bound tensors can be allocated before physical
|
# committed watermark) so upper-bound tensors can be allocated before physical
|
||||||
@@ -110,34 +77,21 @@ class KvVmmArena:
|
|||||||
# (they race and one loads a half-relinked copy -> undefined symbol crash).
|
# (they race and one loads a half-relinked copy -> undefined symbol crash).
|
||||||
self._sfx = f"{os.getpid()}_{KvVmmArena._instance_count}"
|
self._sfx = f"{os.getpid()}_{KvVmmArena._instance_count}"
|
||||||
KvVmmArena._instance_count += 1
|
KvVmmArena._instance_count += 1
|
||||||
drv = _driver()
|
|
||||||
with torch.cuda.device(self.device_id):
|
with torch.cuda.device(self.device_id):
|
||||||
_check(drv.cuInit(0), "cuInit")
|
prop = make_device_allocation_prop(self.device_id)
|
||||||
self._prop = drv.CUmemAllocationProp()
|
self.handle_type = prop.requestedHandleTypes
|
||||||
self._prop.type = drv.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
|
self.granularity = get_device_granularity(self.device_id)
|
||||||
self._prop.location.type = drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
|
||||||
self._prop.location.id = self.device_id
|
|
||||||
self.granularity = query_granularity(self.device_id)
|
|
||||||
self._access = drv.CUmemAccessDesc()
|
|
||||||
self._access.location.type = (
|
|
||||||
drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
|
||||||
)
|
|
||||||
self._access.location.id = self.device_id
|
|
||||||
self._access.flags = (
|
|
||||||
drv.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE
|
|
||||||
)
|
|
||||||
|
|
||||||
self.reserved = self._align(reserve_bytes)
|
self.reserved = self._align(reserve_bytes)
|
||||||
# Align the base to granularity so base + (granularity-aligned cursor) is
|
# Align the base to granularity so base + (granularity-aligned cursor) is
|
||||||
# always a valid cuMemMap address for per-buffer commit_range().
|
# always a valid cuMemMap address for per-buffer commit_range().
|
||||||
self.base = int(
|
self._allocation = VmmReservation(
|
||||||
_check(
|
self.reserved,
|
||||||
drv.cuMemAddressReserve(self.reserved, self.granularity, 0, 0),
|
prop,
|
||||||
"cuMemAddressReserve",
|
self.device_id,
|
||||||
|
alignment=self.granularity,
|
||||||
)
|
)
|
||||||
)
|
self.base = self._allocation.base
|
||||||
# commit_range bookkeeping: mapped VA -> (size, handle); committed bytes per offset.
|
|
||||||
self._ranges = {}
|
|
||||||
self._committed_by_offset = {}
|
self._committed_by_offset = {}
|
||||||
self._range_backed = 0
|
self._range_backed = 0
|
||||||
self._closed = False
|
self._closed = False
|
||||||
@@ -153,11 +107,12 @@ class KvVmmArena:
|
|||||||
self.pool = torch.cuda.MemPool(self._allocator, no_split=True)
|
self.pool = torch.cuda.MemPool(self._allocator, no_split=True)
|
||||||
logger.info(
|
logger.info(
|
||||||
"KvVmmArena[%s] ready: device=%d reserved_va=%.1f GiB "
|
"KvVmmArena[%s] ready: device=%d reserved_va=%.1f GiB "
|
||||||
"granularity=%d KiB",
|
"granularity=%d KiB handle_type=%s",
|
||||||
self._sfx,
|
self._sfx,
|
||||||
self.device_id,
|
self.device_id,
|
||||||
self.reserved / (1024**3),
|
self.reserved / (1024**3),
|
||||||
self.granularity // 1024,
|
self.granularity // 1024,
|
||||||
|
allocation_handle_type_name(self.handle_type),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _align(self, v: int) -> int:
|
def _align(self, v: int) -> int:
|
||||||
@@ -183,16 +138,16 @@ class KvVmmArena:
|
|||||||
)
|
)
|
||||||
self._so_path = f"{out_dir}/{libname}.so"
|
self._so_path = f"{out_dir}/{libname}.so"
|
||||||
lib = ctypes.CDLL(self._so_path)
|
lib = ctypes.CDLL(self._so_path)
|
||||||
self._fn_set_base = getattr(lib, f"kvarena_set_base_{self._sfx}")
|
self._fn_set_base = lib[f"kvarena_set_base_{self._sfx}"]
|
||||||
self._fn_set_base.argtypes = [ctypes.c_void_p]
|
self._fn_set_base.argtypes = [ctypes.c_void_p]
|
||||||
self._fn_set_base.restype = None
|
self._fn_set_base.restype = None
|
||||||
self._fn_set_reserved = getattr(lib, f"kvarena_set_reserved_{self._sfx}")
|
self._fn_set_reserved = lib[f"kvarena_set_reserved_{self._sfx}"]
|
||||||
self._fn_set_reserved.argtypes = [ctypes.c_size_t]
|
self._fn_set_reserved.argtypes = [ctypes.c_size_t]
|
||||||
self._fn_set_reserved.restype = None
|
self._fn_set_reserved.restype = None
|
||||||
self._fn_set_align = getattr(lib, f"kvarena_set_align_{self._sfx}")
|
self._fn_set_align = lib[f"kvarena_set_align_{self._sfx}"]
|
||||||
self._fn_set_align.argtypes = [ctypes.c_size_t]
|
self._fn_set_align.argtypes = [ctypes.c_size_t]
|
||||||
self._fn_set_align.restype = None
|
self._fn_set_align.restype = None
|
||||||
self._fn_cursor = getattr(lib, f"kvarena_cursor_{self._sfx}")
|
self._fn_cursor = lib[f"kvarena_cursor_{self._sfx}"]
|
||||||
self._fn_cursor.argtypes = []
|
self._fn_cursor.argtypes = []
|
||||||
self._fn_cursor.restype = ctypes.c_size_t
|
self._fn_cursor.restype = ctypes.c_size_t
|
||||||
return lib
|
return lib
|
||||||
@@ -217,24 +172,13 @@ class KvVmmArena:
|
|||||||
f"commit_range [{offset}, {offset + want}) exceeds reservation "
|
f"commit_range [{offset}, {offset + want}) exceeds reservation "
|
||||||
f"{self.reserved}"
|
f"{self.reserved}"
|
||||||
)
|
)
|
||||||
drv = _driver()
|
|
||||||
add = want - prev
|
add = want - prev
|
||||||
addr = self.base + offset + prev
|
|
||||||
with torch.cuda.device(self.device_id):
|
with torch.cuda.device(self.device_id):
|
||||||
handle = _check(drv.cuMemCreate(add, self._prop, 0), "cuMemCreate")
|
self._allocation.map(
|
||||||
try:
|
offset + prev,
|
||||||
_check(drv.cuMemMap(addr, add, 0, handle, 0), "cuMemMap")
|
add,
|
||||||
_check(
|
retain_handle=True,
|
||||||
drv.cuMemSetAccess(addr, add, [self._access], 1), "cuMemSetAccess"
|
|
||||||
)
|
)
|
||||||
except Exception:
|
|
||||||
# Roll back this failed extension; leave already-mapped ranges intact.
|
|
||||||
unmap = drv.cuMemUnmap(addr, add)
|
|
||||||
unmap = unmap[0] if isinstance(unmap, tuple) else unmap
|
|
||||||
rel = drv.cuMemRelease(handle)
|
|
||||||
rel = rel[0] if isinstance(rel, tuple) else rel
|
|
||||||
raise
|
|
||||||
self._ranges[addr] = (add, handle)
|
|
||||||
self._committed_by_offset[offset] = want
|
self._committed_by_offset[offset] = want
|
||||||
self._range_backed += add
|
self._range_backed += add
|
||||||
|
|
||||||
@@ -251,25 +195,11 @@ class KvVmmArena:
|
|||||||
if self._closed:
|
if self._closed:
|
||||||
return
|
return
|
||||||
self._closed = True
|
self._closed = True
|
||||||
drv = _driver()
|
|
||||||
try:
|
try:
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
except Exception as e: # pragma: no cover
|
except Exception as e: # pragma: no cover
|
||||||
logger.warning("KvVmmArena.close synchronize failed: %s", e)
|
logger.warning("KvVmmArena.close synchronize failed: %s", e)
|
||||||
for addr, (size, handle) in self._ranges.items():
|
self._allocation.close()
|
||||||
err = drv.cuMemUnmap(addr, size)
|
|
||||||
err = err[0] if isinstance(err, tuple) else err
|
|
||||||
if err != drv.CUresult.CUDA_SUCCESS:
|
|
||||||
logger.warning("cuMemUnmap range -> %s", err)
|
|
||||||
err = drv.cuMemRelease(handle)
|
|
||||||
err = err[0] if isinstance(err, tuple) else err
|
|
||||||
if err != drv.CUresult.CUDA_SUCCESS:
|
|
||||||
logger.warning("cuMemRelease range -> %s", err)
|
|
||||||
self._ranges.clear()
|
|
||||||
err = drv.cuMemAddressFree(self.base, self.reserved)
|
|
||||||
err = err[0] if isinstance(err, tuple) else err
|
|
||||||
if err != drv.CUresult.CUDA_SUCCESS:
|
|
||||||
logger.warning("cuMemAddressFree -> %s", err)
|
|
||||||
|
|
||||||
|
|
||||||
# torch's caching allocator hands the pluggable allocator whole large-pool segments
|
# torch's caching allocator hands the pluggable allocator whole large-pool segments
|
||||||
@@ -326,12 +256,11 @@ class KvVmmBufferOwner:
|
|||||||
|
|
||||||
itemsize = store_dtype.itemsize
|
itemsize = store_dtype.itemsize
|
||||||
with torch.cuda.device(self.device_id):
|
with torch.cuda.device(self.device_id):
|
||||||
gran = query_granularity(self.device_id)
|
gran = get_device_granularity(self.device_id)
|
||||||
reserved_spans = [d.reserved_span_bytes(itemsize) for d in buffer_descs]
|
reserved_spans = [d.reserved_span_bytes(itemsize) for d in buffer_descs]
|
||||||
aligned = [align_up(s, gran) for s in reserved_spans]
|
aligned = [align_up(s, gran) for s in reserved_spans]
|
||||||
reserve_bytes = sum(a + _PER_BUFFER_VA_SLACK for a in aligned) + gran
|
reserve_bytes = sum(a + _PER_BUFFER_VA_SLACK for a in aligned) + gran
|
||||||
self._arena = KvVmmArena(self.device_id, reserve_bytes=reserve_bytes)
|
self._arena = KvVmmArena(self.device_id, reserve_bytes=reserve_bytes)
|
||||||
assert self._arena.granularity == gran, (self._arena.granularity, gran)
|
|
||||||
|
|
||||||
# NORMAL torch tensors through the arena MemPool; torch.empty never touches
|
# NORMAL torch tensors through the arena MemPool; torch.empty never touches
|
||||||
# the unbacked tail.
|
# the unbacked tail.
|
||||||
|
|||||||
@@ -10,15 +10,21 @@ from dataclasses import dataclass
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.distributed.device_communicators.vmm_utils import (
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
_FD_SEND_TIMEOUT_S,
|
_FD_SEND_TIMEOUT_S,
|
||||||
|
VmmReservation,
|
||||||
_get_cuda_driver,
|
_get_cuda_driver,
|
||||||
_recv_fd,
|
_recv_fd,
|
||||||
_send_fd,
|
_send_fd,
|
||||||
|
align_up,
|
||||||
|
allocation_handle_type_name,
|
||||||
check_drv,
|
check_drv,
|
||||||
|
get_allocation_granularity,
|
||||||
|
get_device_allocation_handle_type,
|
||||||
import_and_map_alloc,
|
import_and_map_alloc,
|
||||||
make_rw_access_desc,
|
make_device_allocation_prop,
|
||||||
release_mappings,
|
release_mappings,
|
||||||
|
tensor_from_pointer,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.schedule_batch import (
|
from sglang.srt.managers.schedule_batch import (
|
||||||
Modality,
|
Modality,
|
||||||
@@ -40,18 +46,6 @@ _CONTROL_ALIGNMENT = 256
|
|||||||
_CONTROL_WORD_BYTES = 4
|
_CONTROL_WORD_BYTES = 4
|
||||||
|
|
||||||
|
|
||||||
def _align_up(value: int, alignment: int) -> int:
|
|
||||||
return (value + alignment - 1) // alignment * alignment
|
|
||||||
|
|
||||||
|
|
||||||
def _tensor_from_pointer(pointer: int, size: int, device_index: int) -> torch.Tensor:
|
|
||||||
device = torch.device(f"cuda:{device_index}")
|
|
||||||
storage = torch._C._construct_storage_from_data_pointer(pointer, device, size)
|
|
||||||
return torch.empty(0, dtype=torch.uint8, device=device).set_(
|
|
||||||
storage, 0, (size,), (1,)
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class _PosixFdBroker:
|
class _PosixFdBroker:
|
||||||
"""Serve one exported CUDA allocation FD to local consumer processes."""
|
"""Serve one exported CUDA allocation FD to local consumer processes."""
|
||||||
|
|
||||||
@@ -143,7 +137,7 @@ def _build_packed_tensor_layout(
|
|||||||
layouts = []
|
layouts = []
|
||||||
next_offset = 0
|
next_offset = 0
|
||||||
for tensor in tensors:
|
for tensor in tensors:
|
||||||
next_offset = _align_up(next_offset, tensor.element_size())
|
next_offset = align_up(next_offset, tensor.element_size())
|
||||||
data_nbytes = tensor.numel() * tensor.element_size()
|
data_nbytes = tensor.numel() * tensor.element_size()
|
||||||
layouts.append(
|
layouts.append(
|
||||||
_CudaVmmPackedTensorLayout(
|
_CudaVmmPackedTensorLayout(
|
||||||
@@ -201,9 +195,7 @@ class CudaVmmMemoryPool:
|
|||||||
self._pool_error: BaseException | None = None
|
self._pool_error: BaseException | None = None
|
||||||
self._closed = False
|
self._closed = False
|
||||||
|
|
||||||
self._allocation_handle = None
|
self._allocation: VmmReservation | None = None
|
||||||
self._pool_pointer = None
|
|
||||||
self._allocation_mapped = False
|
|
||||||
self.allocation_size = 0
|
self.allocation_size = 0
|
||||||
self.shareable_handle = None
|
self.shareable_handle = None
|
||||||
self.memory_pool = None
|
self.memory_pool = None
|
||||||
@@ -212,17 +204,29 @@ class CudaVmmMemoryPool:
|
|||||||
self._recycle_stream = None
|
self._recycle_stream = None
|
||||||
self._recycle_thread = None
|
self._recycle_thread = None
|
||||||
|
|
||||||
self.use_fabric = True
|
drv = _get_cuda_driver()
|
||||||
|
fabric = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
|
||||||
|
posix_fd = (
|
||||||
|
drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
|
||||||
|
)
|
||||||
|
self.handle_type = get_device_allocation_handle_type(self.device_index)
|
||||||
|
if self.handle_type == posix_fd and not allow_posix_fallback:
|
||||||
|
raise RuntimeError(
|
||||||
|
"CUDA VMM multimodal transport selected POSIX_FD, but this "
|
||||||
|
"pool requires FABRIC"
|
||||||
|
)
|
||||||
|
self.use_fabric = self.handle_type == fabric
|
||||||
try:
|
try:
|
||||||
self._allocate(memory_size)
|
self._allocate(memory_size)
|
||||||
except RuntimeError as error:
|
except RuntimeError as error:
|
||||||
if not allow_posix_fallback:
|
if not allow_posix_fallback or self.handle_type != fabric:
|
||||||
raise
|
raise
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"CUDA FABRIC VMM allocation is unavailable; falling back to "
|
"CUDA FABRIC VMM allocation is unavailable; falling back to "
|
||||||
"a POSIX FD handle: %s",
|
"a POSIX FD handle: %s",
|
||||||
error,
|
error,
|
||||||
)
|
)
|
||||||
|
self.handle_type = posix_fd
|
||||||
self.use_fabric = False
|
self.use_fabric = False
|
||||||
self._allocate(memory_size)
|
self._allocate(memory_size)
|
||||||
try:
|
try:
|
||||||
@@ -271,30 +275,14 @@ class CudaVmmMemoryPool:
|
|||||||
|
|
||||||
def _allocate(self, memory_size: int) -> None:
|
def _allocate(self, memory_size: int) -> None:
|
||||||
drv = _get_cuda_driver()
|
drv = _get_cuda_driver()
|
||||||
handle_type = (
|
prop = make_device_allocation_prop(
|
||||||
drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
|
self.device_index,
|
||||||
if self.use_fabric
|
handle_types=self.handle_type,
|
||||||
else drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
|
gpu_direct_rdma=self.use_fabric,
|
||||||
)
|
)
|
||||||
prop = drv.CUmemAllocationProp()
|
|
||||||
prop.type = drv.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
|
|
||||||
prop.location.type = drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
|
||||||
prop.location.id = self.device_index
|
|
||||||
prop.requestedHandleTypes = handle_type
|
|
||||||
if self.use_fabric:
|
|
||||||
prop.allocFlags.gpuDirectRDMACapable = 1
|
|
||||||
|
|
||||||
recommended = (
|
|
||||||
drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED
|
|
||||||
)
|
|
||||||
with torch.cuda.device(self.device_index):
|
with torch.cuda.device(self.device_index):
|
||||||
check_drv(drv.cuInit(0), "cuInit")
|
granularity = get_allocation_granularity(prop)
|
||||||
granularity = int(
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemGetAllocationGranularity(prop, recommended),
|
|
||||||
"cuMemGetAllocationGranularity(VMM transport)",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
allocation_size = memory_size // granularity * granularity
|
allocation_size = memory_size // granularity * granularity
|
||||||
if allocation_size == 0:
|
if allocation_size == 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -302,59 +290,47 @@ class CudaVmmMemoryPool:
|
|||||||
f"granularity={granularity}"
|
f"granularity={granularity}"
|
||||||
)
|
)
|
||||||
|
|
||||||
handle = pointer = exported = None
|
allocation = VmmReservation(
|
||||||
mapped = False
|
allocation_size,
|
||||||
|
prop,
|
||||||
|
self.device_index,
|
||||||
|
alignment=granularity,
|
||||||
|
)
|
||||||
|
exported = None
|
||||||
try:
|
try:
|
||||||
handle = check_drv(
|
handle = allocation.map(
|
||||||
drv.cuMemCreate(allocation_size, prop, 0),
|
0,
|
||||||
"cuMemCreate(VMM transport)",
|
allocation_size,
|
||||||
)
|
retain_handle=True,
|
||||||
pointer = int(
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemAddressReserve(allocation_size, granularity, 0, 0),
|
|
||||||
"cuMemAddressReserve(VMM transport)",
|
|
||||||
)
|
|
||||||
)
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemMap(pointer, allocation_size, 0, handle, 0),
|
|
||||||
"cuMemMap(VMM transport)",
|
|
||||||
)
|
|
||||||
mapped = True
|
|
||||||
access = make_rw_access_desc(self.device_index)
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemSetAccess(pointer, allocation_size, [access], 1),
|
|
||||||
"cuMemSetAccess(VMM transport)",
|
|
||||||
)
|
)
|
||||||
exported = check_drv(
|
exported = check_drv(
|
||||||
drv.cuMemExportToShareableHandle(handle, handle_type, 0),
|
drv.cuMemExportToShareableHandle(handle, self.handle_type, 0),
|
||||||
"cuMemExportToShareableHandle(VMM transport)",
|
"cuMemExportToShareableHandle(VMM transport)",
|
||||||
)
|
)
|
||||||
memory_pool = _tensor_from_pointer(
|
memory_pool = tensor_from_pointer(
|
||||||
pointer, allocation_size, self.device_index
|
allocation.base, allocation_size, device_id=self.device_index
|
||||||
)
|
)
|
||||||
except BaseException:
|
except BaseException:
|
||||||
if mapped:
|
allocation.close()
|
||||||
drv.cuMemUnmap(pointer, allocation_size)
|
|
||||||
if pointer is not None:
|
|
||||||
drv.cuMemAddressFree(pointer, allocation_size)
|
|
||||||
if handle is not None:
|
|
||||||
drv.cuMemRelease(handle)
|
|
||||||
if not self.use_fabric and exported is not None:
|
if not self.use_fabric and exported is not None:
|
||||||
os.close(int(exported))
|
os.close(int(exported))
|
||||||
raise
|
raise
|
||||||
|
|
||||||
self._allocation_handle = handle
|
self._allocation = allocation
|
||||||
self._pool_pointer = pointer
|
|
||||||
self._allocation_mapped = True
|
|
||||||
self.allocation_size = allocation_size
|
self.allocation_size = allocation_size
|
||||||
self.shareable_handle = (
|
self.shareable_handle = (
|
||||||
bytes(exported.data) if self.use_fabric else int(exported)
|
bytes(exported.data) if self.use_fabric else int(exported)
|
||||||
)
|
)
|
||||||
|
logger.info(
|
||||||
|
"CUDA VMM multimodal pool uses %s backing on device %d",
|
||||||
|
allocation_handle_type_name(self.handle_type),
|
||||||
|
self.device_index,
|
||||||
|
)
|
||||||
self.memory_pool = memory_pool
|
self.memory_pool = memory_pool
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def control_size(self) -> int:
|
def control_size(self) -> int:
|
||||||
return _align_up(self.consumer_count * _CONTROL_WORD_BYTES, _CONTROL_ALIGNMENT)
|
return align_up(self.consumer_count * _CONTROL_WORD_BYTES, _CONTROL_ALIGNMENT)
|
||||||
|
|
||||||
def _raise_if_failed(self) -> None:
|
def _raise_if_failed(self) -> None:
|
||||||
if self._pool_error is not None:
|
if self._pool_error is not None:
|
||||||
@@ -383,7 +359,7 @@ class CudaVmmMemoryPool:
|
|||||||
if not tensor.is_contiguous():
|
if not tensor.is_contiguous():
|
||||||
tensor = tensor.contiguous()
|
tensor = tensor.contiguous()
|
||||||
data_nbytes = tensor.numel() * tensor.element_size()
|
data_nbytes = tensor.numel() * tensor.element_size()
|
||||||
required_size = _align_up(self.control_size + data_nbytes, _CONTROL_ALIGNMENT)
|
required_size = align_up(self.control_size + data_nbytes, _CONTROL_ALIGNMENT)
|
||||||
source_bytes = tensor.reshape(-1).view(torch.uint8)
|
source_bytes = tensor.reshape(-1).view(torch.uint8)
|
||||||
|
|
||||||
chunk = self._reserve_for_publish(required_size)
|
chunk = self._reserve_for_publish(required_size)
|
||||||
@@ -454,7 +430,7 @@ class CudaVmmMemoryPool:
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
layouts, packed_data_nbytes = _build_packed_tensor_layout(tensors)
|
layouts, packed_data_nbytes = _build_packed_tensor_layout(tensors)
|
||||||
required_size = _align_up(
|
required_size = align_up(
|
||||||
self.control_size + packed_data_nbytes, _CONTROL_ALIGNMENT
|
self.control_size + packed_data_nbytes, _CONTROL_ALIGNMENT
|
||||||
)
|
)
|
||||||
chunk = self._reserve_for_publish(required_size)
|
chunk = self._reserve_for_publish(required_size)
|
||||||
@@ -626,28 +602,11 @@ class CudaVmmMemoryPool:
|
|||||||
if not self.use_fabric and self.shareable_handle is not None:
|
if not self.use_fabric and self.shareable_handle is not None:
|
||||||
os.close(self.shareable_handle)
|
os.close(self.shareable_handle)
|
||||||
self.shareable_handle = None
|
self.shareable_handle = None
|
||||||
if self._pool_pointer is None and self._allocation_handle is None:
|
if self._allocation is None:
|
||||||
return
|
return
|
||||||
drv = _get_cuda_driver()
|
|
||||||
with torch.cuda.device(self.device_index):
|
with torch.cuda.device(self.device_index):
|
||||||
if self._allocation_mapped:
|
self._allocation.close()
|
||||||
check_drv(
|
self._allocation = None
|
||||||
drv.cuMemUnmap(self._pool_pointer, self.allocation_size),
|
|
||||||
"cuMemUnmap(VMM transport pool)",
|
|
||||||
)
|
|
||||||
self._allocation_mapped = False
|
|
||||||
if self._pool_pointer is not None:
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemAddressFree(self._pool_pointer, self.allocation_size),
|
|
||||||
"cuMemAddressFree(VMM transport pool)",
|
|
||||||
)
|
|
||||||
self._pool_pointer = None
|
|
||||||
if self._allocation_handle is not None:
|
|
||||||
check_drv(
|
|
||||||
drv.cuMemRelease(self._allocation_handle),
|
|
||||||
"cuMemRelease(VMM transport pool)",
|
|
||||||
)
|
|
||||||
self._allocation_handle = None
|
|
||||||
|
|
||||||
def shutdown(self) -> None:
|
def shutdown(self) -> None:
|
||||||
with self._shutdown_lock:
|
with self._shutdown_lock:
|
||||||
@@ -727,8 +686,8 @@ def _get_imported_pool(
|
|||||||
peer_rank=-1,
|
peer_rank=-1,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
memory = _tensor_from_pointer(
|
memory = tensor_from_pointer(
|
||||||
pointer, allocation_size, device_index
|
pointer, allocation_size, device_id=device_index
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
release_mappings(
|
release_mappings(
|
||||||
@@ -1039,8 +998,11 @@ class CudaVmmFeatureTransport:
|
|||||||
updates.append((item, "feature", tensor, proxy))
|
updates.append((item, "feature", tensor, proxy))
|
||||||
|
|
||||||
for item in mm_items:
|
for item in mm_items:
|
||||||
for field in ("feature", "precomputed_embeddings"):
|
fields = (
|
||||||
tensor = getattr(item, field)
|
("feature", item.feature),
|
||||||
|
("precomputed_embeddings", item.precomputed_embeddings),
|
||||||
|
)
|
||||||
|
for field, tensor in fields:
|
||||||
if _contains_tensor_container(tensor):
|
if _contains_tensor_container(tensor):
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
"CUDA VMM feature transport requires each feature "
|
"CUDA VMM feature transport requires each feature "
|
||||||
@@ -1074,8 +1036,11 @@ class CudaVmmFeatureTransport:
|
|||||||
|
|
||||||
errors = []
|
errors = []
|
||||||
for item in mm_items:
|
for item in mm_items:
|
||||||
for field in ("feature", "precomputed_embeddings"):
|
fields = (
|
||||||
proxy = getattr(item, field)
|
("feature", item.feature),
|
||||||
|
("precomputed_embeddings", item.precomputed_embeddings),
|
||||||
|
)
|
||||||
|
for field, proxy in fields:
|
||||||
if not isinstance(proxy, CudaVmmTensorTransportProxy):
|
if not isinstance(proxy, CudaVmmTensorTransportProxy):
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -442,7 +442,7 @@ class TestCudaVmmTransport(CustomTestCase):
|
|||||||
def test_posix_export_fd_closes_when_allocation_setup_fails(self):
|
def test_posix_export_fd_closes_when_allocation_setup_fails(self):
|
||||||
with (
|
with (
|
||||||
patch(
|
patch(
|
||||||
"sglang.srt.utils.cuda_vmm_transport_utils._tensor_from_pointer",
|
"sglang.srt.utils.cuda_vmm_transport_utils.tensor_from_pointer",
|
||||||
side_effect=RuntimeError("forced storage failure"),
|
side_effect=RuntimeError("forced storage failure"),
|
||||||
),
|
),
|
||||||
patch(
|
patch(
|
||||||
|
|||||||
@@ -18,48 +18,27 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
pool.memory_pool = object()
|
pool.memory_pool = object()
|
||||||
pool.use_fabric = True
|
pool.use_fabric = True
|
||||||
pool.shareable_handle = b"handle"
|
pool.shareable_handle = b"handle"
|
||||||
pool._pool_pointer = 123
|
allocation = MagicMock()
|
||||||
pool._allocation_handle = 456
|
allocation.close.side_effect = [
|
||||||
pool._allocation_mapped = True
|
RuntimeError("forced allocation close failure"),
|
||||||
pool.allocation_size = 4096
|
None,
|
||||||
|
]
|
||||||
|
pool._allocation = allocation
|
||||||
pool.device_index = 0
|
pool.device_index = 0
|
||||||
driver = MagicMock()
|
|
||||||
driver.cuMemUnmap.return_value = "unmap"
|
|
||||||
driver.cuMemAddressFree.return_value = "address_free"
|
|
||||||
driver.cuMemRelease.return_value = "release"
|
|
||||||
failed_once = False
|
|
||||||
|
|
||||||
def check_driver(result, _operation):
|
|
||||||
nonlocal failed_once
|
|
||||||
if result == "address_free" and not failed_once:
|
|
||||||
failed_once = True
|
|
||||||
raise RuntimeError("forced address-free failure")
|
|
||||||
return result
|
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch.object(vmm, "_get_cuda_driver", return_value=driver),
|
|
||||||
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
|
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
|
||||||
patch.object(vmm, "check_drv", side_effect=check_driver),
|
self.assertRaisesRegex(RuntimeError, "forced allocation close failure"),
|
||||||
self.assertRaisesRegex(RuntimeError, "forced address-free failure"),
|
|
||||||
):
|
):
|
||||||
pool._release_allocation()
|
pool._release_allocation()
|
||||||
|
|
||||||
self.assertFalse(pool._allocation_mapped)
|
self.assertIs(pool._allocation, allocation)
|
||||||
self.assertEqual(pool._pool_pointer, 123)
|
|
||||||
self.assertEqual(pool._allocation_handle, 456)
|
|
||||||
|
|
||||||
with (
|
with patch.object(vmm.torch.cuda, "device", return_value=nullcontext()):
|
||||||
patch.object(vmm, "_get_cuda_driver", return_value=driver),
|
|
||||||
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
|
|
||||||
patch.object(vmm, "check_drv", side_effect=lambda result, _: result),
|
|
||||||
):
|
|
||||||
pool._release_allocation()
|
pool._release_allocation()
|
||||||
|
|
||||||
self.assertIsNone(pool._pool_pointer)
|
self.assertIsNone(pool._allocation)
|
||||||
self.assertIsNone(pool._allocation_handle)
|
self.assertEqual(allocation.close.call_count, 2)
|
||||||
self.assertEqual(driver.cuMemUnmap.call_count, 1)
|
|
||||||
self.assertEqual(driver.cuMemAddressFree.call_count, 2)
|
|
||||||
self.assertEqual(driver.cuMemRelease.call_count, 1)
|
|
||||||
|
|
||||||
def test_model_class_controls_cuda_vmm_opt_in(self):
|
def test_model_class_controls_cuda_vmm_opt_in(self):
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
|||||||
+66
-1
@@ -21,11 +21,15 @@ import torch.distributed as dist
|
|||||||
from cuda.bindings import driver as drv
|
from cuda.bindings import driver as drv
|
||||||
|
|
||||||
from sglang.kernels.jit.utils import cache_once
|
from sglang.kernels.jit.utils import cache_once
|
||||||
from sglang.srt.distributed.device_communicators.vmm_utils import (
|
from sglang.srt import cuda_vmm_utils
|
||||||
|
from sglang.srt.cuda_vmm_utils import (
|
||||||
check_drv,
|
check_drv,
|
||||||
exchange_posix_fds,
|
exchange_posix_fds,
|
||||||
export_shareable_handles,
|
export_shareable_handles,
|
||||||
|
get_allocation_granularity,
|
||||||
|
get_device_allocation_handle_type,
|
||||||
import_and_map_alloc,
|
import_and_map_alloc,
|
||||||
|
make_device_allocation_prop,
|
||||||
make_rw_access_desc,
|
make_rw_access_desc,
|
||||||
map_chunk_into_span,
|
map_chunk_into_span,
|
||||||
release_mappings,
|
release_mappings,
|
||||||
@@ -107,6 +111,67 @@ def _assert_region(va: int, expected: int, peer: int, chunk: int) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("rejected", "expected"),
|
||||||
|
[
|
||||||
|
((_FABRIC,), _POSIX_FD),
|
||||||
|
((_FABRIC, _POSIX_FD), 0),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_default_handle_type_fallback(monkeypatch, rejected, expected) -> None:
|
||||||
|
device_id = torch.cuda.current_device()
|
||||||
|
create = drv.cuMemCreate
|
||||||
|
|
||||||
|
def reject_selected(size, prop, flags):
|
||||||
|
if prop.requestedHandleTypes in rejected:
|
||||||
|
return (drv.CUresult.CUDA_ERROR_NOT_SUPPORTED, None)
|
||||||
|
return create(size, prop, flags)
|
||||||
|
|
||||||
|
get_device_allocation_handle_type.cache_clear()
|
||||||
|
monkeypatch.setattr(cuda_vmm_utils, "is_gpu_fabric_ready", lambda _device: True)
|
||||||
|
monkeypatch.setattr(drv, "cuMemCreate", reject_selected)
|
||||||
|
try:
|
||||||
|
selected = get_device_allocation_handle_type(device_id)
|
||||||
|
prop = make_device_allocation_prop(device_id)
|
||||||
|
assert selected == expected
|
||||||
|
assert prop.requestedHandleTypes == expected
|
||||||
|
assert prop.allocFlags.gpuDirectRDMACapable == 0
|
||||||
|
|
||||||
|
explicit = make_device_allocation_prop(
|
||||||
|
device_id,
|
||||||
|
handle_types=_FABRIC,
|
||||||
|
gpu_direct_rdma=True,
|
||||||
|
)
|
||||||
|
assert explicit.requestedHandleTypes == _FABRIC
|
||||||
|
assert explicit.allocFlags.gpuDirectRDMACapable == 1
|
||||||
|
|
||||||
|
non_exportable = make_device_allocation_prop(device_id, handle_types=None)
|
||||||
|
assert non_exportable.requestedHandleTypes == 0
|
||||||
|
|
||||||
|
explicit_none = make_device_allocation_prop(device_id, handle_types=0)
|
||||||
|
assert explicit_none.requestedHandleTypes == 0
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="handle_types must be"):
|
||||||
|
make_device_allocation_prop(device_id, handle_types="fabric")
|
||||||
|
with pytest.raises(ValueError, match="invalid CUDA handle-type value"):
|
||||||
|
make_device_allocation_prop(device_id, handle_types=42)
|
||||||
|
finally:
|
||||||
|
get_device_allocation_handle_type.cache_clear()
|
||||||
|
|
||||||
|
|
||||||
|
def test_granularity_defaults_to_recommended(monkeypatch) -> None:
|
||||||
|
prop = make_device_allocation_prop(0, handle_types=None)
|
||||||
|
seen = []
|
||||||
|
|
||||||
|
def granularity(_prop, flag):
|
||||||
|
seen.append(flag)
|
||||||
|
return (drv.CUresult.CUDA_SUCCESS, _ALLOC_BYTES)
|
||||||
|
|
||||||
|
monkeypatch.setattr(drv, "cuMemGetAllocationGranularity", granularity)
|
||||||
|
assert get_allocation_granularity(prop) == _ALLOC_BYTES
|
||||||
|
assert seen == [_RECOMMENDED]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("n_chunks", [1, 3])
|
@pytest.mark.parametrize("n_chunks", [1, 3])
|
||||||
@pytest.mark.parametrize("transport", ["posix", "fabric"])
|
@pytest.mark.parametrize("transport", ["posix", "fabric"])
|
||||||
def test_handle_roundtrip(transport: str, n_chunks: int) -> None:
|
def test_handle_roundtrip(transport: str, n_chunks: int) -> None:
|
||||||
Reference in New Issue
Block a user