From 5ae214f18bc4d2231433b21863283e83e8f313a8 Mon Sep 17 00:00:00 2001 From: Ming Yang Date: Wed, 1 Jul 2026 14:02:00 -0700 Subject: [PATCH] Extract reusable VMM shareable-handle helpers from register_graph_inputs (#29621) --- .../custom_all_reduce_v2.py | 2 +- ...m_all_reduce_vmm_utils.py => vmm_utils.py} | 565 ++++++++++-------- .../unit/distributed/test_vmm_utils.py | 209 +++++++ 3 files changed, 530 insertions(+), 246 deletions(-) rename python/sglang/srt/distributed/device_communicators/{custom_all_reduce_vmm_utils.py => vmm_utils.py} (54%) create mode 100644 test/registered/unit/distributed/test_vmm_utils.py diff --git a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py b/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py index 2cbc5b0fe..a98707724 100644 --- a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py +++ b/python/sglang/srt/distributed/device_communicators/custom_all_reduce_v2.py @@ -12,7 +12,7 @@ from sglang.srt.distributed.device_communicators.custom_all_reduce_utils import can_use_custom_all_reduce_with_nvlink, is_weak_contiguous, ) -from sglang.srt.distributed.device_communicators.custom_all_reduce_vmm_utils import ( +from sglang.srt.distributed.device_communicators.vmm_utils import ( VmmGraphInputManager, is_vmm_pointer, ) diff --git a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_vmm_utils.py b/python/sglang/srt/distributed/device_communicators/vmm_utils.py similarity index 54% rename from python/sglang/srt/distributed/device_communicators/custom_all_reduce_vmm_utils.py rename to python/sglang/srt/distributed/device_communicators/vmm_utils.py index 239f6de68..895bd259c 100644 --- a/python/sglang/srt/distributed/device_communicators/custom_all_reduce_vmm_utils.py +++ b/python/sglang/srt/distributed/device_communicators/vmm_utils.py @@ -27,7 +27,7 @@ def _get_cuda_driver(): return _drv -def _check_drv(result_tuple, label): +def check_drv(result_tuple, label): """Check a cuda.bindings driver call result and return the value.""" if not isinstance(result_tuple, tuple): result_tuple = (result_tuple,) @@ -52,6 +52,36 @@ def is_vmm_pointer(ptr: int) -> bool: return False +def make_rw_access_desc(device_id: int): + """A read-write, device-local ``CUmemAccessDesc`` for ``device_id``.""" + drv = _get_cuda_driver() + desc = drv.CUmemAccessDesc() + desc.location.type = drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE + desc.location.id = device_id + desc.flags = drv.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE + return desc + + +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) + dist.all_reduce(flag, op=dist.ReduceOp.BAND, group=group) + return flag.item() == 1 + + +def release_mappings(mappings) -> None: + """Unmap + address-free each ``(va, span_size, [(rel, size), ...])`` mapping. + + Pops from ``mappings`` so a partially-released list is safe to retry. + """ + drv = _get_cuda_driver() + while mappings: + va, span_size, mapped_chunks = mappings.pop() + for rel, size in mapped_chunks: + check_drv(drv.cuMemUnmap(int(va) + int(rel), int(size)), "cuMemUnmap") + check_drv(drv.cuMemAddressFree(int(va), int(span_size)), "cuMemAddressFree") + + def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None: import array import socket @@ -94,6 +124,259 @@ def _recv_fd(sock): return int(src_rank), int(base_idx), int(fds[0]) +def export_shareable_handles(retained_handles, group: ProcessGroup, rank: int): + """Export retained VMM handles, preferring FABRIC and falling back to POSIX fds. + + FABRIC is used only if every rank can export it; otherwise all ranks use POSIX + fds. Returns ``(fabric_handles, posix_fds, use_fabric)`` (one list populated); + raises if both fail on any rank. Caller owns the returned ``posix_fds``. + """ + drv = _get_cuda_driver() + FABRIC = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC + POSIX_FD = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR + + fabric_handles: List[bytes] = [] + fabric_error: Optional[Exception] = None + try: + for alloc_h in retained_handles: + fabric_h = check_drv( + drv.cuMemExportToShareableHandle(alloc_h, FABRIC, 0), + "cuMemExportToShareableHandle(FABRIC)", + ) + fabric_handles.append(bytes(fabric_h.data)) + fabric_ok = True + except Exception as e: + fabric_error = e + fabric_ok = False + fabric_handles = [] + logger.info( + "FABRIC handle export failed on rank %s; falling back to " + "POSIX fd transport: %s", + rank, + e, + ) + + if all_ranks_ok(group, fabric_ok): + return fabric_handles, [], True + + posix_fds: List[int] = [] + posix_error: Optional[Exception] = None + try: + for alloc_h in retained_handles: + fd = check_drv( + drv.cuMemExportToShareableHandle(alloc_h, POSIX_FD, 0), + "cuMemExportToShareableHandle(POSIX_FD)", + ) + posix_fds.append(int(fd)) + posix_ok = True + except Exception as e: + posix_error = e + posix_ok = False + for fd in posix_fds: + try: + os.close(fd) + except OSError: + pass + posix_fds = [] + + if not all_ranks_ok(group, posix_ok): + cause = posix_error or fabric_error + message = ( + "VMM handle export failed: FABRIC export failed on at least one " + "rank and POSIX fd export failed on at least one rank" + ) + if cause is not None: + message += f"; local rank {rank} error: {cause}" + raise RuntimeError(message) from posix_error + + return [], posix_fds, False + + +def exchange_posix_fds( + group: ProcessGroup, + rank: int, + world_size: int, + local_fds: List[int], + peer_base_counts: List[int], +): + """Exchange POSIX file descriptors across ranks via SCM_RIGHTS over a UNIX + 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_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) + server.settimeout(_FD_SEND_TIMEOUT_S) + received_fds = {} + errors = [] + + def recv_loop(): + try: + for _ in range(world_size - 1): + conn, _ = server.accept() + with conn: + conn.settimeout(_FD_SEND_TIMEOUT_S) + while True: + packet = _recv_fd(conn) + if packet is None: + break + src_rank, base_idx, fd = packet + key = (src_rank, base_idx) + if key in received_fds: + os.close(fd) + raise RuntimeError(f"duplicate fd for {key}") + received_fds[key] = fd + except BaseException as e: + errors.append(e) + + try: + server.bind(sock_path) + server.listen(world_size) + paths = [None] * world_size + dist.all_gather_object(paths, sock_path, group=group) + + thread = threading.Thread(target=recv_loop, daemon=True) + thread.start() + try: + for peer_rank, peer_path in enumerate(paths): + if peer_rank == rank: + continue + with socket.socket(socket.AF_UNIX, sock_kind) as sock: + sock.settimeout(_FD_SEND_TIMEOUT_S) + sock.connect(peer_path) + for base_idx, fd in enumerate(local_fds): + _send_fd(sock, fd, rank, base_idx) + finally: + thread.join(_FD_SEND_TIMEOUT_S) + + if thread.is_alive(): + raise RuntimeError("timed out waiting for POSIX fd exchange") + if errors: + raise RuntimeError("POSIX fd exchange receive failed") from errors[0] + + expected = { + (src_rank, base_idx) + for src_rank, count in enumerate(peer_base_counts) + if src_rank != rank + for base_idx in range(count) + } + missing = expected.difference(received_fds) + extra = set(received_fds).difference(expected) + if missing or extra: + for fd in received_fds.values(): + os.close(fd) + raise RuntimeError( + "POSIX fd exchange mismatch: " + f"missing={sorted(missing)[:8]}, extra={sorted(extra)[:8]}" + ) + return received_fds + finally: + server.close() + try: + os.unlink(sock_path) + except FileNotFoundError: + pass + try: + os.rmdir(sock_dir) + except OSError: + pass + + +def import_peer_handle(fabric_handle, fd, *, use_fabric: bool, peer_rank: int): + """Import a peer allocation handle (FABRIC or POSIX fd). Returns the handle. + + For POSIX the fd is duped before import so the caller keeps ownership of the + original. + """ + drv = _get_cuda_driver() + if use_fabric: + FABRIC = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC + return check_drv( + drv.cuMemImportFromShareableHandle(fabric_handle, FABRIC), + f"cuMemImportFromShareableHandle(rank={peer_rank})", + ) + POSIX_FD = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR + dup_fd = os.dup(fd) + try: + return check_drv( + drv.cuMemImportFromShareableHandle(dup_fd, POSIX_FD), + f"cuMemImportFromShareableHandle(rank={peer_rank}, POSIX_FD)", + ) + finally: + try: + os.close(dup_fd) + except OSError: + pass + + +def import_and_map_alloc( + fabric_handle, + fd, + alloc_size: int, + device_id: int, + *, + use_fabric: bool, + peer_rank: int, +) -> int: + """Import a peer allocation, map it at a freshly reserved VA, return the VA.""" + drv = _get_cuda_driver() + imp_h = import_peer_handle( + fabric_handle, fd, use_fabric=use_fabric, peer_rank=peer_rank + ) + prop = check_drv( + drv.cuMemGetAllocationPropertiesFromHandle(imp_h), + "cuMemGetAllocationPropertiesFromHandle", + ) + gran = check_drv( + drv.cuMemGetAllocationGranularity( + prop, + drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED, + ), + "cuMemGetAllocationGranularity", + ) + va = check_drv( + drv.cuMemAddressReserve(alloc_size, int(gran), 0, 0), "cuMemAddressReserve" + ) + check_drv(drv.cuMemMap(int(va), alloc_size, 0, imp_h, 0), "cuMemMap") + access = make_rw_access_desc(device_id) + check_drv(drv.cuMemSetAccess(int(va), alloc_size, [access], 1), "cuMemSetAccess") + check_drv(drv.cuMemRelease(imp_h), "cuMemRelease(peer)") + return int(va) + + +def map_chunk_into_span( + fabric_handle, + fd, + span_va: int, + rel: int, + alloc_size: int, + device_id: int, + *, + use_fabric: bool, + peer_rank: int, +) -> None: + """Import + map a peer chunk into a caller-reserved span at ``span_va + rel``.""" + drv = _get_cuda_driver() + imp_h = import_peer_handle( + fabric_handle, fd, use_fabric=use_fabric, peer_rank=peer_rank + ) + check_drv( + drv.cuMemMap(int(span_va) + rel, int(alloc_size), 0, imp_h, 0), + "cuMemMap(span)", + ) + access = make_rw_access_desc(device_id) + check_drv( + drv.cuMemSetAccess(int(span_va) + rel, int(alloc_size), [access], 1), + "cuMemSetAccess(span)", + ) + check_drv(drv.cuMemRelease(imp_h), "cuMemRelease(span)") + + class VmmGraphInputManager: def __init__( self, @@ -117,11 +400,6 @@ class VmmGraphInputManager: allocations, then registers the peer VAs. FABRIC handles are preferred; POSIX file descriptors are used when FABRIC is unavailable. """ - drv = _get_cuda_driver() - FABRIC = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC - POSIX_FD = ( - drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR - ) FABRIC_HANDLE_BYTES = 64 MAX_VMM_BASES = 4096 MAX_CHUNKS_PER_INPUT = 16 @@ -142,68 +420,20 @@ class VmmGraphInputManager: f"Too many VMM bases to share: {num_bases} > {MAX_VMM_BASES}" ) - local_fabric_handles: List[bytes] = [] + drv = _get_cuda_driver() local_posix_fds: List[int] = [] retained_handles = [] try: for base_ptr, _ in bases_info: - alloc_h = _check_drv( + alloc_h = check_drv( drv.cuMemRetainAllocationHandle(base_ptr), "cuMemRetainAllocationHandle", ) retained_handles.append(alloc_h) - local_fabric_error: Optional[Exception] = None - try: - for alloc_h in retained_handles: - fabric_h = _check_drv( - drv.cuMemExportToShareableHandle(alloc_h, FABRIC, 0), - "cuMemExportToShareableHandle(FABRIC)", - ) - local_fabric_handles.append(bytes(fabric_h.data)) - local_fabric_ok = True - except Exception as e: - local_fabric_error = e - local_fabric_ok = False - local_fabric_handles = [] - logger.info( - "FABRIC handle export failed on rank %s; falling back to " - "POSIX fd transport: %s", - self.rank, - e, - ) - - use_fabric = self._all_ranks_ok(local_fabric_ok) - if not use_fabric: - local_posix_error: Optional[Exception] = None - try: - for alloc_h in retained_handles: - fd = _check_drv( - drv.cuMemExportToShareableHandle(alloc_h, POSIX_FD, 0), - "cuMemExportToShareableHandle(POSIX_FD)", - ) - local_posix_fds.append(int(fd)) - local_posix_ok = True - except Exception as e: - local_posix_error = e - local_posix_ok = False - for fd in local_posix_fds: - try: - os.close(fd) - except OSError: - pass - local_posix_fds = [] - - if not self._all_ranks_ok(local_posix_ok): - local_cause = local_posix_error or local_fabric_error - message = ( - "VMM graph input registration failed: FABRIC export " - "failed on at least one rank and POSIX fd export failed " - "on at least one rank" - ) - if local_cause is not None: - message += f"; local rank {self.rank} error: {local_cause}" - raise RuntimeError(message) from local_posix_error + local_fabric_handles, local_posix_fds, use_fabric = ( + export_shareable_handles(retained_handles, self.group, self.rank) + ) local_input_chunks = [ [int(idx) for idx in indices] for indices in input_chunk_indices @@ -304,7 +534,10 @@ class VmmGraphInputManager: posix_peer_fds = {} if not use_fabric: - posix_peer_fds = self._exchange_posix_fds( + posix_peer_fds = exchange_posix_fds( + self.group, + self.rank, + self.world_size, local_posix_fds, [len(peer_bases) for peer_bases in all_base_payload], ) @@ -316,25 +549,6 @@ class VmmGraphInputManager: peer_span_va = {} # (rank, chunk_indices...) -> (local VA, peer base) new_mappings = [] - def import_peer_handle(peer_rank: int, base_idx: int, fabric_handle): - if use_fabric: - return _check_drv( - drv.cuMemImportFromShareableHandle(fabric_handle, FABRIC), - f"cuMemImportFromShareableHandle(rank={peer_rank})", - ) - fd = posix_peer_fds[(peer_rank, base_idx)] - dup_fd = os.dup(fd) - try: - return _check_drv( - drv.cuMemImportFromShareableHandle(dup_fd, POSIX_FD), - f"cuMemImportFromShareableHandle(rank={peer_rank}, POSIX_FD)", - ) - finally: - try: - os.close(dup_fd) - except OSError: - pass - try: for peer_rank in range(self.world_size): if peer_rank == self.rank: @@ -344,41 +558,17 @@ class VmmGraphInputManager: peer_bases = all_base_payload[peer_rank] for idx, (_, fb, alloc_size) in enumerate(peer_bases): - imp_h = import_peer_handle(peer_rank, idx, fb) - prop = _check_drv( - drv.cuMemGetAllocationPropertiesFromHandle(imp_h), - "cuMemGetAllocationPropertiesFromHandle", + fd = None if use_fabric else posix_peer_fds[(peer_rank, idx)] + va = import_and_map_alloc( + fb, + fd, + alloc_size, + device_id, + use_fabric=use_fabric, + peer_rank=peer_rank, ) - gran = _check_drv( - drv.cuMemGetAllocationGranularity( - prop, - drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED, - ), - "cuMemGetAllocationGranularity", - ) - va = _check_drv( - drv.cuMemAddressReserve(alloc_size, int(gran), 0, 0), - "cuMemAddressReserve", - ) - _check_drv( - drv.cuMemMap(int(va), alloc_size, 0, imp_h, 0), - "cuMemMap", - ) - access = drv.CUmemAccessDesc() - access.location.type = ( - drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE - ) - access.location.id = device_id - access.flags = ( - drv.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE - ) - _check_drv( - drv.cuMemSetAccess(int(va), alloc_size, [access], 1), - "cuMemSetAccess", - ) - peer_base_va[(peer_rank, idx)] = int(va) - new_mappings.append((int(va), alloc_size, [(0, alloc_size)])) - _check_drv(drv.cuMemRelease(imp_h), "cuMemRelease(peer)") + peer_base_va[(peer_rank, idx)] = va + new_mappings.append((va, alloc_size, [(0, alloc_size)])) # Build per-input peer VA lists and register. peer_ptrs = [] @@ -402,7 +592,7 @@ class VmmGraphInputManager: if rank == self.rank: span_va = int(first_base) else: - span_va = _check_drv( + span_va = check_drv( drv.cuMemAddressReserve(span_size, 0, 0, 0), "cuMemAddressReserve(span)", ) @@ -410,38 +600,22 @@ class VmmGraphInputManager: for chunk_idx in chunks: base_ptr, fb, alloc_size = peer_bases[chunk_idx] rel = int(base_ptr) - int(first_base) - imp_h = import_peer_handle(rank, chunk_idx, fb) - _check_drv( - drv.cuMemMap( - int(span_va) + rel, - int(alloc_size), - 0, - imp_h, - 0, - ), - "cuMemMap(span)", + fd = ( + None + if use_fabric + else posix_peer_fds[(rank, chunk_idx)] ) - access = drv.CUmemAccessDesc() - access.location.type = ( - drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE - ) - access.location.id = device_id - access.flags = ( - drv.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE - ) - _check_drv( - drv.cuMemSetAccess( - int(span_va) + rel, - int(alloc_size), - [access], - 1, - ), - "cuMemSetAccess(span)", + map_chunk_into_span( + fb, + fd, + span_va, + rel, + int(alloc_size), + device_id, + use_fabric=use_fabric, + peer_rank=rank, ) mapped_chunks.append((rel, int(alloc_size))) - _check_drv( - drv.cuMemRelease(imp_h), "cuMemRelease(span)" - ) new_mappings.append( (int(span_va), span_size, mapped_chunks) ) @@ -454,7 +628,7 @@ class VmmGraphInputManager: self.obj.register_peer_mapped_inputs(peer_ptrs) self._peer_mappings.extend(new_mappings) except Exception: - self._release_peer_mappings(new_mappings) + release_mappings(new_mappings) raise finally: for fd in posix_peer_fds.values(): @@ -472,108 +646,9 @@ class VmmGraphInputManager: for fd in local_posix_fds: os.close(fd) for h in retained_handles: - _check_drv(drv.cuMemRelease(h), "cuMemRelease(retained)") + check_drv(drv.cuMemRelease(h), "cuMemRelease(retained)") def close(self): if not self._peer_mappings: return - self._release_peer_mappings(self._peer_mappings) - - def _all_ranks_ok(self, ok: bool) -> bool: - flag = torch.tensor([1 if ok else 0], dtype=torch.int32) - dist.all_reduce(flag, op=dist.ReduceOp.BAND, group=self.group) - return flag.item() == 1 - - def _exchange_posix_fds(self, local_fds: List[int], peer_base_counts: List[int]): - import socket - import tempfile - import threading - - sock_kind = getattr(socket, "SOCK_SEQPACKET", socket.SOCK_STREAM) - sock_dir = tempfile.mkdtemp(prefix="sgl_ar_fd_") - sock_path = os.path.join(sock_dir, f"rank_{self.rank}.sock") - server = socket.socket(socket.AF_UNIX, sock_kind) - server.settimeout(_FD_SEND_TIMEOUT_S) - received_fds = {} - errors = [] - - def recv_loop(): - try: - for _ in range(self.world_size - 1): - conn, _ = server.accept() - with conn: - conn.settimeout(_FD_SEND_TIMEOUT_S) - while True: - packet = _recv_fd(conn) - if packet is None: - break - src_rank, base_idx, fd = packet - key = (src_rank, base_idx) - if key in received_fds: - os.close(fd) - raise RuntimeError(f"duplicate fd for {key}") - received_fds[key] = fd - except BaseException as e: - errors.append(e) - - try: - server.bind(sock_path) - server.listen(self.world_size) - paths = [None] * self.world_size - dist.all_gather_object(paths, sock_path, group=self.group) - - thread = threading.Thread(target=recv_loop, daemon=True) - thread.start() - try: - for peer_rank, peer_path in enumerate(paths): - if peer_rank == self.rank: - continue - with socket.socket(socket.AF_UNIX, sock_kind) as sock: - sock.settimeout(_FD_SEND_TIMEOUT_S) - sock.connect(peer_path) - for base_idx, fd in enumerate(local_fds): - _send_fd(sock, fd, self.rank, base_idx) - finally: - thread.join(_FD_SEND_TIMEOUT_S) - - if thread.is_alive(): - raise RuntimeError("timed out waiting for POSIX fd exchange") - if errors: - raise RuntimeError("POSIX fd exchange receive failed") from errors[0] - - expected = { - (rank, base_idx) - for rank, count in enumerate(peer_base_counts) - if rank != self.rank - for base_idx in range(count) - } - missing = expected.difference(received_fds) - extra = set(received_fds).difference(expected) - if missing or extra: - for fd in received_fds.values(): - os.close(fd) - raise RuntimeError( - "POSIX fd exchange mismatch: " - f"missing={sorted(missing)[:8]}, extra={sorted(extra)[:8]}" - ) - return received_fds - finally: - server.close() - try: - os.unlink(sock_path) - except FileNotFoundError: - pass - try: - os.rmdir(sock_dir) - except OSError: - pass - - def _release_peer_mappings(self, mappings): - drv = _get_cuda_driver() - while mappings: - va, span_size, mapped_chunks = mappings.pop() - for rel, size in mapped_chunks: - _check_drv(drv.cuMemUnmap(int(va) + int(rel), int(size)), "cuMemUnmap") - _check_drv( - drv.cuMemAddressFree(int(va), int(span_size)), "cuMemAddressFree" - ) + release_mappings(self._peer_mappings) diff --git a/test/registered/unit/distributed/test_vmm_utils.py b/test/registered/unit/distributed/test_vmm_utils.py new file mode 100644 index 000000000..8637969cf --- /dev/null +++ b/test/registered/unit/distributed/test_vmm_utils.py @@ -0,0 +1,209 @@ +"""Unit tests for the VMM cross-process handle helpers in ``vmm_utils``. + +Round-trips export -> exchange -> import/map across ranks for both transports +(POSIX, FABRIC) and both mapping shapes (single base, multi-chunk span). The +only in-tree consumer, ``register_graph_inputs``, reaches this path only under +``expandable_segments``, so the tests allocate shareable buffers directly. A +POSIX-only allocation forces ``export_shareable_handles`` down its POSIX +fallback (otherwise unreachable on FABRIC hardware); FABRIC cases need an +NVLink fabric (GB200/GB300) and skip elsewhere. +""" + +from __future__ import annotations + +import atexit +import os + +import numpy as np +import pytest +import torch +import torch.distributed as dist +from cuda.bindings import driver as drv + +from sglang.jit_kernel.tests.utils import multigpu_pytest_main +from sglang.jit_kernel.utils import cache_once +from sglang.srt.distributed.device_communicators.vmm_utils import ( + check_drv, + exchange_posix_fds, + export_shareable_handles, + import_and_map_alloc, + make_rw_access_desc, + map_chunk_into_span, + release_mappings, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=60, stage="base-b", runner_config="2-gpu-large") + +_FABRIC = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_FABRIC +_POSIX_FD = drv.CUmemAllocationHandleType.CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR +_RECOMMENDED = drv.CUmemAllocationGranularity_flags.CU_MEM_ALLOC_GRANULARITY_RECOMMENDED +_ALLOC_BYTES = 2 * 1024 * 1024 + + +@cache_once +def _gloo_group() -> dist.ProcessGroup: + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + dist.init_process_group(backend="gloo") + atexit.register(dist.destroy_process_group) + return dist.group.WORLD + + +def _make_prop(handle_type, device_id: int): + 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 = device_id + prop.requestedHandleTypes = handle_type + return prop + + +@cache_once +def _fabric_available() -> bool: + """True if this device can create + export FABRIC handles (GB200/GB300).""" + prop = _make_prop(_FABRIC, torch.cuda.current_device()) + err, gran = drv.cuMemGetAllocationGranularity(prop, _RECOMMENDED) + if err != drv.CUresult.CUDA_SUCCESS: + return False + err, handle = drv.cuMemCreate(int(gran), prop, 0) + if err != drv.CUresult.CUDA_SUCCESS: + return False + err, _ = drv.cuMemExportToShareableHandle(handle, _FABRIC, 0) + drv.cuMemRelease(handle) + return err == drv.CUresult.CUDA_SUCCESS + + +def _create_alloc(handle_type, size_hint: int): + """Create a mapped, RW, shareable VMM allocation. Returns (handle, va, size).""" + device_id = torch.cuda.current_device() + prop = _make_prop(handle_type, device_id) + gran = check_drv( + drv.cuMemGetAllocationGranularity(prop, _RECOMMENDED), + "cuMemGetAllocationGranularity", + ) + size = ((size_hint + gran - 1) // gran) * gran + handle = check_drv(drv.cuMemCreate(size, prop, 0), "cuMemCreate") + va = check_drv(drv.cuMemAddressReserve(size, gran, 0, 0), "cuMemAddressReserve") + check_drv(drv.cuMemMap(int(va), size, 0, handle, 0), "cuMemMap") + check_drv( + drv.cuMemSetAccess(int(va), size, [make_rw_access_desc(device_id)], 1), + "cuMemSetAccess", + ) + return handle, int(va), size + + +def _byte(rank: int, chunk: int) -> int: + """A distinct nonzero fill byte per (rank, chunk).""" + return (rank * 16 + chunk + 1) & 0xFF + + +def _assert_region(va: int, expected: int, peer: int, chunk: int) -> None: + host = np.empty(16, dtype=np.uint8) + check_drv(drv.cuMemcpyDtoH(host.ctypes.data, va, host.nbytes), "cuMemcpyDtoH") + assert ( + host == expected + ).all(), ( + f"read {host.tolist()} from peer {peer} chunk {chunk}, expected all {expected}" + ) + + +@pytest.mark.parametrize("n_chunks", [1, 3]) +@pytest.mark.parametrize("transport", ["posix", "fabric"]) +def test_handle_roundtrip(transport: str, n_chunks: int) -> None: + group = _gloo_group() + if transport == "fabric" and not _fabric_available(): + pytest.skip("FABRIC handles require an NVLink fabric (GB200/GB300)") + rank = dist.get_rank(group) + world = dist.get_world_size(group) + device_id = torch.cuda.current_device() + handle_type = _FABRIC if transport == "fabric" else _POSIX_FD + + handles, vas, sizes = [], [], [] + for chunk in range(n_chunks): + handle, va, size = _create_alloc(handle_type, _ALLOC_BYTES) + check_drv(drv.cuMemsetD8(va, _byte(rank, chunk), size), "cuMemsetD8") + handles.append(handle) + vas.append(va) + sizes.append(size) + torch.cuda.synchronize() + + posix_fds, peer_fds, mappings = [], {}, [] + try: + fabric_handles, posix_fds, use_fabric = export_shareable_handles( + handles, group, rank + ) + assert use_fabric == (transport == "fabric") + + # FABRIC handles travel inline; POSIX fds are exchanged out-of-band + # (process-local). + local_meta = [ + (sizes[c], fabric_handles[c] if use_fabric else None) + for c in range(n_chunks) + ] + all_meta = [None] * world + dist.all_gather_object(all_meta, local_meta, group=group) + if not use_fabric: + peer_fds = exchange_posix_fds( + group, rank, world, posix_fds, [n_chunks] * world + ) + + for peer in range(world): + if peer == rank: + continue + peer_meta = all_meta[peer] + if n_chunks == 1: + size, fabric_handle = peer_meta[0] + fd = None if use_fabric else peer_fds[(peer, 0)] + peer_va = import_and_map_alloc( + fabric_handle, + fd, + size, + device_id, + use_fabric=use_fabric, + peer_rank=peer, + ) + mappings.append((peer_va, size, [(0, size)])) + _assert_region(peer_va, _byte(peer, 0), peer, 0) + continue + + span_size = sum(size for size, _ in peer_meta) + span_va = int( + check_drv( + drv.cuMemAddressReserve(span_size, 0, 0, 0), + "cuMemAddressReserve(span)", + ) + ) + rel, mapped = 0, [] + for chunk, (size, fabric_handle) in enumerate(peer_meta): + fd = None if use_fabric else peer_fds[(peer, chunk)] + map_chunk_into_span( + fabric_handle, + fd, + span_va, + rel, + size, + device_id, + use_fabric=use_fabric, + peer_rank=peer, + ) + mapped.append((rel, size)) + rel += size + mappings.append((span_va, span_size, mapped)) + rel = 0 + for chunk, (size, _) in enumerate(peer_meta): + _assert_region(span_va + rel, _byte(peer, chunk), peer, chunk) + rel += size + finally: + release_mappings(mappings) + for fd in peer_fds.values(): + os.close(fd) + for fd in posix_fds: + os.close(fd) + for handle, va, size in zip(handles, vas, sizes): + check_drv(drv.cuMemUnmap(va, size), "cuMemUnmap") + check_drv(drv.cuMemAddressFree(va, size), "cuMemAddressFree") + check_drv(drv.cuMemRelease(handle), "cuMemRelease") + + +if __name__ == "__main__": + multigpu_pytest_main(__name__, __file__, num_gpus=(2,))