Consolidate CUDA VMM allocation helpers (#34199)
This commit is contained in:
+380
-26
@@ -1,7 +1,13 @@
|
||||
import array
|
||||
import ctypes
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import struct
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from functools import cache
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import torch
|
||||
@@ -12,18 +18,43 @@ from sglang.srt.utils import log_info_on_rank0
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_drv = None
|
||||
_FD_HEADER_BYTES = 24
|
||||
_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():
|
||||
"""Lazily import cuda.bindings.driver (cached after first call)."""
|
||||
global _drv
|
||||
"""Return the imported CUDA driver bindings."""
|
||||
if _drv is None:
|
||||
from cuda.bindings import driver
|
||||
|
||||
_drv = driver
|
||||
raise ImportError("cuda.bindings.driver is required for CUDA VMM operations")
|
||||
return _drv
|
||||
|
||||
|
||||
@@ -38,6 +69,22 @@ def check_drv(result_tuple, label):
|
||||
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:
|
||||
"""Check if a device pointer is VMM-backed (cuMemCreate/cuMemMap).
|
||||
|
||||
@@ -112,6 +159,329 @@ def make_rw_access_desc(device_id: int):
|
||||
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:
|
||||
"""True iff ``ok`` holds on every rank in ``group`` (BAND all-reduce)."""
|
||||
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:
|
||||
import array
|
||||
import socket
|
||||
|
||||
fds = array.array("i", [int(fd)])
|
||||
header = struct.pack("<QQQ", int(src_rank), int(base_idx), 1)
|
||||
sent = sock.sendmsg(
|
||||
@@ -147,9 +514,6 @@ def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None:
|
||||
|
||||
|
||||
def _recv_fd(sock):
|
||||
import array
|
||||
import socket
|
||||
|
||||
fd_item_size = array.array("i").itemsize
|
||||
data, ancdata, _, _ = sock.recvmsg(
|
||||
_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
|
||||
owns the received fds and must close them.
|
||||
"""
|
||||
import socket
|
||||
import tempfile
|
||||
import threading
|
||||
|
||||
sock_kind = getattr(socket, "SOCK_SEQPACKET", socket.SOCK_STREAM)
|
||||
sock_kind = socket.SOCK_SEQPACKET
|
||||
sock_dir = tempfile.mkdtemp(prefix="sgl_ar_fd_")
|
||||
sock_path = os.path.join(sock_dir, f"rank_{rank}.sock")
|
||||
server = socket.socket(socket.AF_UNIX, sock_kind)
|
||||
@@ -382,13 +742,7 @@ def import_and_map_alloc(
|
||||
drv.cuMemGetAllocationPropertiesFromHandle(imp_h),
|
||||
"cuMemGetAllocationPropertiesFromHandle",
|
||||
)
|
||||
gran = check_drv(
|
||||
drv.cuMemGetAllocationGranularity(
|
||||
prop,
|
||||
drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED,
|
||||
),
|
||||
"cuMemGetAllocationGranularity",
|
||||
)
|
||||
gran = get_allocation_granularity(prop)
|
||||
va = check_drv(
|
||||
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
|
||||
graph capture pointers into unique base allocations via cuMemGetAddressRange.
|
||||
Python exports handles for each unique base, imports + cuMemMaps peer
|
||||
allocations, then registers the peer VAs. FABRIC handles are preferred;
|
||||
POSIX file descriptors are used when FABRIC is unavailable.
|
||||
allocations, then registers the peer virtual addresses. FABRIC handles are
|
||||
preferred; POSIX file descriptors are used when FABRIC is unavailable.
|
||||
"""
|
||||
FABRIC_HANDLE_BYTES = 64
|
||||
MAX_VMM_BASES = 4096
|
||||
@@ -19,6 +19,7 @@ import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
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.parallel_state import in_the_same_node_as
|
||||
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
|
||||
|
||||
|
||||
# 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
|
||||
def is_one_nvlink_clique(
|
||||
group: torch.distributed.ProcessGroup, device: torch.device
|
||||
|
||||
@@ -31,6 +31,11 @@ from sglang.kernels.ops.communication.all_reduce import (
|
||||
IPCManager,
|
||||
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.environ import envs
|
||||
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_weak_contiguous,
|
||||
)
|
||||
from .vmm_utils import (
|
||||
VmmGraphInputManager,
|
||||
compute_graph_capture_bases,
|
||||
is_vmm_pointer,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Dict, List, Optional, Tuple
|
||||
|
||||
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
|
||||
PeerRanges = List[Tuple[int, int]]
|
||||
|
||||
@@ -4,14 +4,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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,
|
||||
create_local_handle,
|
||||
get_allocation_granularity,
|
||||
map_handle,
|
||||
release_handle,
|
||||
check_drv,
|
||||
get_device_granularity,
|
||||
make_device_allocation_prop,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -31,7 +33,8 @@ class PagePool:
|
||||
page_size: Optional[int] = None,
|
||||
):
|
||||
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:
|
||||
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):
|
||||
handles = []
|
||||
for _ in range(num_pages):
|
||||
h = create_local_handle(self._page_size, device_id)
|
||||
handles.append(h)
|
||||
reservation = VmmReservation(
|
||||
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)
|
||||
logger.debug(
|
||||
f"PagePool slot {slot_idx}: {num_pages} pages × {self._page_size} B"
|
||||
@@ -78,21 +88,21 @@ class PagePool:
|
||||
def map_pages(
|
||||
self,
|
||||
slot: int,
|
||||
va_start: int,
|
||||
reservation: VmmReservation,
|
||||
offset: 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
|
||||
) -> None:
|
||||
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
|
||||
reservation.map_existing(
|
||||
offset + i * self._page_size,
|
||||
self._page_size,
|
||||
handle,
|
||||
)
|
||||
|
||||
def release(self) -> None:
|
||||
if self._released:
|
||||
@@ -100,7 +110,7 @@ class PagePool:
|
||||
self._released = True
|
||||
for handles in self._page_handles:
|
||||
for h in handles:
|
||||
release_handle(h)
|
||||
check_drv(cuda.cuMemRelease(h), "cuMemRelease")
|
||||
self._page_handles = [[], []]
|
||||
|
||||
|
||||
|
||||
@@ -11,30 +11,23 @@ import torch
|
||||
import torch.distributed as dist
|
||||
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,
|
||||
exchange_posix_fds,
|
||||
export_shareable_handles,
|
||||
get_device_granularity,
|
||||
import_peer_handle,
|
||||
make_device_allocation_prop,
|
||||
tensor_from_pointer,
|
||||
)
|
||||
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__)
|
||||
|
||||
@@ -54,7 +47,8 @@ def _copy_local_weights_to_handles(
|
||||
layout: DwdpExpertLayout,
|
||||
device_id: 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] = {}
|
||||
sizes: Dict[Tuple[int, str], int] = {}
|
||||
|
||||
@@ -69,21 +63,17 @@ def _copy_local_weights_to_handles(
|
||||
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)
|
||||
reservation = VmmReservation(phys_size, prop, device_id, alignment=granularity)
|
||||
handle = int(reservation.map(0, phys_size, retain_handle=True))
|
||||
|
||||
nbytes = param.numel() * param.element_size()
|
||||
check_drv(
|
||||
cuda.cuMemcpyDtoD(temp_va + data_offset, param.data_ptr(), nbytes),
|
||||
cuda.cuMemcpyDtoD(reservation.base + data_offset, param.data_ptr(), nbytes),
|
||||
"cuMemcpyDtoD",
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
unmap_va(temp_va, phys_size)
|
||||
free_va(temp_va, phys_size)
|
||||
reservation.close(release_handles=False)
|
||||
|
||||
param.untyped_storage().resize_(0)
|
||||
|
||||
@@ -105,7 +95,7 @@ class DWDPTransport:
|
||||
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]] = []
|
||||
self._peer_reservations: List[VmmReservation] = []
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
@@ -144,7 +134,8 @@ class DWDPTransport:
|
||||
device_id: int,
|
||||
) -> None:
|
||||
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]
|
||||
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_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))
|
||||
peer_reservation = VmmReservation(
|
||||
peer_phys_size,
|
||||
prop,
|
||||
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
|
||||
peer_tensor = tensor_from_ptr(
|
||||
ptr=peer_va + peer_data_offset,
|
||||
peer_tensor = tensor_from_pointer(
|
||||
peer_reservation.base + peer_data_offset,
|
||||
peer_end_bytes - peer_start_bytes,
|
||||
shape=(num_peer_experts,) + spec.full_shape[1:],
|
||||
dtype=spec.dtype,
|
||||
device_id=device_id,
|
||||
@@ -226,13 +222,12 @@ class DWDPTransport:
|
||||
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 reservation in self._peer_reservations:
|
||||
reservation.close()
|
||||
self._peer_reservations.clear()
|
||||
|
||||
for h in self._imported_handles:
|
||||
release_handle(h)
|
||||
check_drv(cuda.cuMemRelease(h), "cuMemRelease")
|
||||
self._imported_handles.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
|
||||
|
||||
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 (
|
||||
EdgeInfo,
|
||||
LayerWeightSpecs,
|
||||
@@ -15,15 +21,6 @@ from sglang.srt.layers.moe.dwdp.layout import (
|
||||
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__)
|
||||
|
||||
@@ -44,7 +41,8 @@ class WeightBuffer:
|
||||
self._local_end = local_end
|
||||
self._dwdp_size = dwdp_size
|
||||
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._page_pool: Optional[PagePool] = None
|
||||
self._moe_layer_indices = sorted(layer_weight_specs.keys())
|
||||
@@ -53,8 +51,7 @@ class WeightBuffer:
|
||||
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._reservations: Dict[int, List[VmmReservation]] = {}
|
||||
self._released = False
|
||||
|
||||
@classmethod
|
||||
@@ -105,8 +102,7 @@ class WeightBuffer:
|
||||
|
||||
self._tensors[layer_idx] = {}
|
||||
self._remote_slices[layer_idx] = {}
|
||||
self._mappings[layer_idx] = []
|
||||
self._va_regions[layer_idx] = []
|
||||
self._reservations[layer_idx] = []
|
||||
|
||||
page_pool_offset = 0
|
||||
|
||||
@@ -114,40 +110,41 @@ class WeightBuffer:
|
||||
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]
|
||||
reservation = VmmReservation(
|
||||
layout.total_size,
|
||||
self._prop,
|
||||
self._device_id,
|
||||
alignment=self._granularity,
|
||||
)
|
||||
self._reservations[layer_idx].append(reservation)
|
||||
va_base = reservation.base
|
||||
|
||||
if layout.pre_size > 0:
|
||||
pre_maps = self._page_pool.map_pages(
|
||||
self._page_pool.map_pages(
|
||||
slot=buf_slot,
|
||||
va_start=va_base,
|
||||
reservation=reservation,
|
||||
offset=0,
|
||||
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))
|
||||
reservation.map_existing(layout.pre_size, layout.mnnvl_size, handle)
|
||||
|
||||
if layout.post_size > 0:
|
||||
post_va = mnnvl_va + layout.mnnvl_size
|
||||
post_maps = self._page_pool.map_pages(
|
||||
self._page_pool.map_pages(
|
||||
slot=buf_slot,
|
||||
va_start=post_va,
|
||||
reservation=reservation,
|
||||
offset=layout.pre_size + layout.mnnvl_size,
|
||||
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,
|
||||
full_tensor = tensor_from_pointer(
|
||||
tensor_start,
|
||||
layout.num_experts * layout.expert_bytes,
|
||||
shape=spec.full_shape,
|
||||
dtype=spec.dtype,
|
||||
device_id=self._device_id,
|
||||
@@ -206,14 +203,10 @@ class WeightBuffer:
|
||||
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()
|
||||
for reservations in self._reservations.values():
|
||||
for reservation in reservations:
|
||||
reservation.close()
|
||||
self._reservations.clear()
|
||||
self._tensors.clear()
|
||||
self._remote_slices.clear()
|
||||
if self._page_pool is not None:
|
||||
|
||||
@@ -11,52 +11,19 @@ import torch
|
||||
import torch.utils.cpp_extension
|
||||
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:
|
||||
from sglang.srt.mem_cache.memory_pool import KvBufferDesc
|
||||
|
||||
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
|
||||
# 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).
|
||||
self._sfx = f"{os.getpid()}_{KvVmmArena._instance_count}"
|
||||
KvVmmArena._instance_count += 1
|
||||
drv = _driver()
|
||||
with torch.cuda.device(self.device_id):
|
||||
_check(drv.cuInit(0), "cuInit")
|
||||
self._prop = drv.CUmemAllocationProp()
|
||||
self._prop.type = drv.CUmemAllocationType.CU_MEM_ALLOCATION_TYPE_PINNED
|
||||
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
|
||||
)
|
||||
prop = make_device_allocation_prop(self.device_id)
|
||||
self.handle_type = prop.requestedHandleTypes
|
||||
self.granularity = get_device_granularity(self.device_id)
|
||||
|
||||
self.reserved = self._align(reserve_bytes)
|
||||
# Align the base to granularity so base + (granularity-aligned cursor) is
|
||||
# always a valid cuMemMap address for per-buffer commit_range().
|
||||
self.base = int(
|
||||
_check(
|
||||
drv.cuMemAddressReserve(self.reserved, self.granularity, 0, 0),
|
||||
"cuMemAddressReserve",
|
||||
)
|
||||
self._allocation = VmmReservation(
|
||||
self.reserved,
|
||||
prop,
|
||||
self.device_id,
|
||||
alignment=self.granularity,
|
||||
)
|
||||
# commit_range bookkeeping: mapped VA -> (size, handle); committed bytes per offset.
|
||||
self._ranges = {}
|
||||
self.base = self._allocation.base
|
||||
self._committed_by_offset = {}
|
||||
self._range_backed = 0
|
||||
self._closed = False
|
||||
@@ -153,11 +107,12 @@ class KvVmmArena:
|
||||
self.pool = torch.cuda.MemPool(self._allocator, no_split=True)
|
||||
logger.info(
|
||||
"KvVmmArena[%s] ready: device=%d reserved_va=%.1f GiB "
|
||||
"granularity=%d KiB",
|
||||
"granularity=%d KiB handle_type=%s",
|
||||
self._sfx,
|
||||
self.device_id,
|
||||
self.reserved / (1024**3),
|
||||
self.granularity // 1024,
|
||||
allocation_handle_type_name(self.handle_type),
|
||||
)
|
||||
|
||||
def _align(self, v: int) -> int:
|
||||
@@ -183,16 +138,16 @@ class KvVmmArena:
|
||||
)
|
||||
self._so_path = f"{out_dir}/{libname}.so"
|
||||
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.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.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.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.restype = ctypes.c_size_t
|
||||
return lib
|
||||
@@ -217,24 +172,13 @@ class KvVmmArena:
|
||||
f"commit_range [{offset}, {offset + want}) exceeds reservation "
|
||||
f"{self.reserved}"
|
||||
)
|
||||
drv = _driver()
|
||||
add = want - prev
|
||||
addr = self.base + offset + prev
|
||||
with torch.cuda.device(self.device_id):
|
||||
handle = _check(drv.cuMemCreate(add, self._prop, 0), "cuMemCreate")
|
||||
try:
|
||||
_check(drv.cuMemMap(addr, add, 0, handle, 0), "cuMemMap")
|
||||
_check(
|
||||
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._allocation.map(
|
||||
offset + prev,
|
||||
add,
|
||||
retain_handle=True,
|
||||
)
|
||||
self._committed_by_offset[offset] = want
|
||||
self._range_backed += add
|
||||
|
||||
@@ -251,25 +195,11 @@ class KvVmmArena:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
drv = _driver()
|
||||
try:
|
||||
torch.cuda.synchronize()
|
||||
except Exception as e: # pragma: no cover
|
||||
logger.warning("KvVmmArena.close synchronize failed: %s", e)
|
||||
for addr, (size, handle) in self._ranges.items():
|
||||
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)
|
||||
self._allocation.close()
|
||||
|
||||
|
||||
# torch's caching allocator hands the pluggable allocator whole large-pool segments
|
||||
@@ -326,12 +256,11 @@ class KvVmmBufferOwner:
|
||||
|
||||
itemsize = store_dtype.itemsize
|
||||
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]
|
||||
aligned = [align_up(s, gran) for s in reserved_spans]
|
||||
reserve_bytes = sum(a + _PER_BUFFER_VA_SLACK for a in aligned) + gran
|
||||
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
|
||||
# the unbacked tail.
|
||||
|
||||
@@ -10,15 +10,21 @@ from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.distributed.device_communicators.vmm_utils import (
|
||||
from sglang.srt.cuda_vmm_utils import (
|
||||
_FD_SEND_TIMEOUT_S,
|
||||
VmmReservation,
|
||||
_get_cuda_driver,
|
||||
_recv_fd,
|
||||
_send_fd,
|
||||
align_up,
|
||||
allocation_handle_type_name,
|
||||
check_drv,
|
||||
get_allocation_granularity,
|
||||
get_device_allocation_handle_type,
|
||||
import_and_map_alloc,
|
||||
make_rw_access_desc,
|
||||
make_device_allocation_prop,
|
||||
release_mappings,
|
||||
tensor_from_pointer,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
Modality,
|
||||
@@ -40,18 +46,6 @@ _CONTROL_ALIGNMENT = 256
|
||||
_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:
|
||||
"""Serve one exported CUDA allocation FD to local consumer processes."""
|
||||
|
||||
@@ -143,7 +137,7 @@ def _build_packed_tensor_layout(
|
||||
layouts = []
|
||||
next_offset = 0
|
||||
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()
|
||||
layouts.append(
|
||||
_CudaVmmPackedTensorLayout(
|
||||
@@ -201,9 +195,7 @@ class CudaVmmMemoryPool:
|
||||
self._pool_error: BaseException | None = None
|
||||
self._closed = False
|
||||
|
||||
self._allocation_handle = None
|
||||
self._pool_pointer = None
|
||||
self._allocation_mapped = False
|
||||
self._allocation: VmmReservation | None = None
|
||||
self.allocation_size = 0
|
||||
self.shareable_handle = None
|
||||
self.memory_pool = None
|
||||
@@ -212,17 +204,29 @@ class CudaVmmMemoryPool:
|
||||
self._recycle_stream = 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:
|
||||
self._allocate(memory_size)
|
||||
except RuntimeError as error:
|
||||
if not allow_posix_fallback:
|
||||
if not allow_posix_fallback or self.handle_type != fabric:
|
||||
raise
|
||||
logger.warning(
|
||||
"CUDA FABRIC VMM allocation is unavailable; falling back to "
|
||||
"a POSIX FD handle: %s",
|
||||
error,
|
||||
)
|
||||
self.handle_type = posix_fd
|
||||
self.use_fabric = False
|
||||
self._allocate(memory_size)
|
||||
try:
|
||||
@@ -271,30 +275,14 @@ class CudaVmmMemoryPool:
|
||||
|
||||
def _allocate(self, memory_size: int) -> None:
|
||||
drv = _get_cuda_driver()
|
||||
handle_type = (
|
||||
drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC
|
||||
if self.use_fabric
|
||||
else drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR
|
||||
prop = make_device_allocation_prop(
|
||||
self.device_index,
|
||||
handle_types=self.handle_type,
|
||||
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):
|
||||
check_drv(drv.cuInit(0), "cuInit")
|
||||
granularity = int(
|
||||
check_drv(
|
||||
drv.cuMemGetAllocationGranularity(prop, recommended),
|
||||
"cuMemGetAllocationGranularity(VMM transport)",
|
||||
)
|
||||
)
|
||||
granularity = get_allocation_granularity(prop)
|
||||
allocation_size = memory_size // granularity * granularity
|
||||
if allocation_size == 0:
|
||||
raise ValueError(
|
||||
@@ -302,59 +290,47 @@ class CudaVmmMemoryPool:
|
||||
f"granularity={granularity}"
|
||||
)
|
||||
|
||||
handle = pointer = exported = None
|
||||
mapped = False
|
||||
allocation = VmmReservation(
|
||||
allocation_size,
|
||||
prop,
|
||||
self.device_index,
|
||||
alignment=granularity,
|
||||
)
|
||||
exported = None
|
||||
try:
|
||||
handle = check_drv(
|
||||
drv.cuMemCreate(allocation_size, prop, 0),
|
||||
"cuMemCreate(VMM transport)",
|
||||
)
|
||||
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)",
|
||||
handle = allocation.map(
|
||||
0,
|
||||
allocation_size,
|
||||
retain_handle=True,
|
||||
)
|
||||
exported = check_drv(
|
||||
drv.cuMemExportToShareableHandle(handle, handle_type, 0),
|
||||
drv.cuMemExportToShareableHandle(handle, self.handle_type, 0),
|
||||
"cuMemExportToShareableHandle(VMM transport)",
|
||||
)
|
||||
memory_pool = _tensor_from_pointer(
|
||||
pointer, allocation_size, self.device_index
|
||||
memory_pool = tensor_from_pointer(
|
||||
allocation.base, allocation_size, device_id=self.device_index
|
||||
)
|
||||
except BaseException:
|
||||
if mapped:
|
||||
drv.cuMemUnmap(pointer, allocation_size)
|
||||
if pointer is not None:
|
||||
drv.cuMemAddressFree(pointer, allocation_size)
|
||||
if handle is not None:
|
||||
drv.cuMemRelease(handle)
|
||||
allocation.close()
|
||||
if not self.use_fabric and exported is not None:
|
||||
os.close(int(exported))
|
||||
raise
|
||||
|
||||
self._allocation_handle = handle
|
||||
self._pool_pointer = pointer
|
||||
self._allocation_mapped = True
|
||||
self._allocation = allocation
|
||||
self.allocation_size = allocation_size
|
||||
self.shareable_handle = (
|
||||
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
|
||||
|
||||
@property
|
||||
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:
|
||||
if self._pool_error is not None:
|
||||
@@ -383,7 +359,7 @@ class CudaVmmMemoryPool:
|
||||
if not tensor.is_contiguous():
|
||||
tensor = tensor.contiguous()
|
||||
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)
|
||||
|
||||
chunk = self._reserve_for_publish(required_size)
|
||||
@@ -454,7 +430,7 @@ class CudaVmmMemoryPool:
|
||||
return []
|
||||
|
||||
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
|
||||
)
|
||||
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:
|
||||
os.close(self.shareable_handle)
|
||||
self.shareable_handle = None
|
||||
if self._pool_pointer is None and self._allocation_handle is None:
|
||||
if self._allocation is None:
|
||||
return
|
||||
drv = _get_cuda_driver()
|
||||
with torch.cuda.device(self.device_index):
|
||||
if self._allocation_mapped:
|
||||
check_drv(
|
||||
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
|
||||
self._allocation.close()
|
||||
self._allocation = None
|
||||
|
||||
def shutdown(self) -> None:
|
||||
with self._shutdown_lock:
|
||||
@@ -727,8 +686,8 @@ def _get_imported_pool(
|
||||
peer_rank=-1,
|
||||
)
|
||||
try:
|
||||
memory = _tensor_from_pointer(
|
||||
pointer, allocation_size, device_index
|
||||
memory = tensor_from_pointer(
|
||||
pointer, allocation_size, device_id=device_index
|
||||
)
|
||||
except Exception:
|
||||
release_mappings(
|
||||
@@ -1039,8 +998,11 @@ class CudaVmmFeatureTransport:
|
||||
updates.append((item, "feature", tensor, proxy))
|
||||
|
||||
for item in mm_items:
|
||||
for field in ("feature", "precomputed_embeddings"):
|
||||
tensor = getattr(item, field)
|
||||
fields = (
|
||||
("feature", item.feature),
|
||||
("precomputed_embeddings", item.precomputed_embeddings),
|
||||
)
|
||||
for field, tensor in fields:
|
||||
if _contains_tensor_container(tensor):
|
||||
raise TypeError(
|
||||
"CUDA VMM feature transport requires each feature "
|
||||
@@ -1074,8 +1036,11 @@ class CudaVmmFeatureTransport:
|
||||
|
||||
errors = []
|
||||
for item in mm_items:
|
||||
for field in ("feature", "precomputed_embeddings"):
|
||||
proxy = getattr(item, field)
|
||||
fields = (
|
||||
("feature", item.feature),
|
||||
("precomputed_embeddings", item.precomputed_embeddings),
|
||||
)
|
||||
for field, proxy in fields:
|
||||
if not isinstance(proxy, CudaVmmTensorTransportProxy):
|
||||
continue
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user