Extract reusable VMM shareable-handle helpers from register_graph_inputs (#29621)
This commit is contained in:
@@ -12,7 +12,7 @@ from sglang.srt.distributed.device_communicators.custom_all_reduce_utils import
|
|||||||
can_use_custom_all_reduce_with_nvlink,
|
can_use_custom_all_reduce_with_nvlink,
|
||||||
is_weak_contiguous,
|
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,
|
VmmGraphInputManager,
|
||||||
is_vmm_pointer,
|
is_vmm_pointer,
|
||||||
)
|
)
|
||||||
|
|||||||
+320
-245
@@ -27,7 +27,7 @@ def _get_cuda_driver():
|
|||||||
return _drv
|
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."""
|
"""Check a cuda.bindings driver call result and return the value."""
|
||||||
if not isinstance(result_tuple, tuple):
|
if not isinstance(result_tuple, tuple):
|
||||||
result_tuple = (result_tuple,)
|
result_tuple = (result_tuple,)
|
||||||
@@ -52,6 +52,36 @@ def is_vmm_pointer(ptr: int) -> bool:
|
|||||||
return False
|
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:
|
def _send_fd(sock, fd: int, src_rank: int, base_idx: int) -> None:
|
||||||
import array
|
import array
|
||||||
import socket
|
import socket
|
||||||
@@ -94,6 +124,259 @@ def _recv_fd(sock):
|
|||||||
return int(src_rank), int(base_idx), int(fds[0])
|
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:
|
class VmmGraphInputManager:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -117,11 +400,6 @@ class VmmGraphInputManager:
|
|||||||
allocations, then registers the peer VAs. FABRIC handles are preferred;
|
allocations, then registers the peer VAs. FABRIC handles are preferred;
|
||||||
POSIX file descriptors are used when FABRIC is unavailable.
|
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
|
FABRIC_HANDLE_BYTES = 64
|
||||||
MAX_VMM_BASES = 4096
|
MAX_VMM_BASES = 4096
|
||||||
MAX_CHUNKS_PER_INPUT = 16
|
MAX_CHUNKS_PER_INPUT = 16
|
||||||
@@ -142,68 +420,20 @@ class VmmGraphInputManager:
|
|||||||
f"Too many VMM bases to share: {num_bases} > {MAX_VMM_BASES}"
|
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] = []
|
local_posix_fds: List[int] = []
|
||||||
retained_handles = []
|
retained_handles = []
|
||||||
try:
|
try:
|
||||||
for base_ptr, _ in bases_info:
|
for base_ptr, _ in bases_info:
|
||||||
alloc_h = _check_drv(
|
alloc_h = check_drv(
|
||||||
drv.cuMemRetainAllocationHandle(base_ptr),
|
drv.cuMemRetainAllocationHandle(base_ptr),
|
||||||
"cuMemRetainAllocationHandle",
|
"cuMemRetainAllocationHandle",
|
||||||
)
|
)
|
||||||
retained_handles.append(alloc_h)
|
retained_handles.append(alloc_h)
|
||||||
|
|
||||||
local_fabric_error: Optional[Exception] = None
|
local_fabric_handles, local_posix_fds, use_fabric = (
|
||||||
try:
|
export_shareable_handles(retained_handles, self.group, self.rank)
|
||||||
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_input_chunks = [
|
local_input_chunks = [
|
||||||
[int(idx) for idx in indices] for indices in input_chunk_indices
|
[int(idx) for idx in indices] for indices in input_chunk_indices
|
||||||
@@ -304,7 +534,10 @@ class VmmGraphInputManager:
|
|||||||
|
|
||||||
posix_peer_fds = {}
|
posix_peer_fds = {}
|
||||||
if not use_fabric:
|
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,
|
local_posix_fds,
|
||||||
[len(peer_bases) for peer_bases in all_base_payload],
|
[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)
|
peer_span_va = {} # (rank, chunk_indices...) -> (local VA, peer base)
|
||||||
new_mappings = []
|
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:
|
try:
|
||||||
for peer_rank in range(self.world_size):
|
for peer_rank in range(self.world_size):
|
||||||
if peer_rank == self.rank:
|
if peer_rank == self.rank:
|
||||||
@@ -344,41 +558,17 @@ class VmmGraphInputManager:
|
|||||||
|
|
||||||
peer_bases = all_base_payload[peer_rank]
|
peer_bases = all_base_payload[peer_rank]
|
||||||
for idx, (_, fb, alloc_size) in enumerate(peer_bases):
|
for idx, (_, fb, alloc_size) in enumerate(peer_bases):
|
||||||
imp_h = import_peer_handle(peer_rank, idx, fb)
|
fd = None if use_fabric else posix_peer_fds[(peer_rank, idx)]
|
||||||
prop = _check_drv(
|
va = import_and_map_alloc(
|
||||||
drv.cuMemGetAllocationPropertiesFromHandle(imp_h),
|
fb,
|
||||||
"cuMemGetAllocationPropertiesFromHandle",
|
fd,
|
||||||
|
alloc_size,
|
||||||
|
device_id,
|
||||||
|
use_fabric=use_fabric,
|
||||||
|
peer_rank=peer_rank,
|
||||||
)
|
)
|
||||||
gran = _check_drv(
|
peer_base_va[(peer_rank, idx)] = va
|
||||||
drv.cuMemGetAllocationGranularity(
|
new_mappings.append((va, alloc_size, [(0, alloc_size)]))
|
||||||
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)")
|
|
||||||
|
|
||||||
# Build per-input peer VA lists and register.
|
# Build per-input peer VA lists and register.
|
||||||
peer_ptrs = []
|
peer_ptrs = []
|
||||||
@@ -402,7 +592,7 @@ class VmmGraphInputManager:
|
|||||||
if rank == self.rank:
|
if rank == self.rank:
|
||||||
span_va = int(first_base)
|
span_va = int(first_base)
|
||||||
else:
|
else:
|
||||||
span_va = _check_drv(
|
span_va = check_drv(
|
||||||
drv.cuMemAddressReserve(span_size, 0, 0, 0),
|
drv.cuMemAddressReserve(span_size, 0, 0, 0),
|
||||||
"cuMemAddressReserve(span)",
|
"cuMemAddressReserve(span)",
|
||||||
)
|
)
|
||||||
@@ -410,38 +600,22 @@ class VmmGraphInputManager:
|
|||||||
for chunk_idx in chunks:
|
for chunk_idx in chunks:
|
||||||
base_ptr, fb, alloc_size = peer_bases[chunk_idx]
|
base_ptr, fb, alloc_size = peer_bases[chunk_idx]
|
||||||
rel = int(base_ptr) - int(first_base)
|
rel = int(base_ptr) - int(first_base)
|
||||||
imp_h = import_peer_handle(rank, chunk_idx, fb)
|
fd = (
|
||||||
_check_drv(
|
None
|
||||||
drv.cuMemMap(
|
if use_fabric
|
||||||
int(span_va) + rel,
|
else posix_peer_fds[(rank, chunk_idx)]
|
||||||
int(alloc_size),
|
|
||||||
0,
|
|
||||||
imp_h,
|
|
||||||
0,
|
|
||||||
),
|
|
||||||
"cuMemMap(span)",
|
|
||||||
)
|
)
|
||||||
access = drv.CUmemAccessDesc()
|
map_chunk_into_span(
|
||||||
access.location.type = (
|
fb,
|
||||||
drv.CUmemLocationType.CU_MEM_LOCATION_TYPE_DEVICE
|
fd,
|
||||||
)
|
span_va,
|
||||||
access.location.id = device_id
|
rel,
|
||||||
access.flags = (
|
int(alloc_size),
|
||||||
drv.CUmemAccess_flags.CU_MEM_ACCESS_FLAGS_PROT_READWRITE
|
device_id,
|
||||||
)
|
use_fabric=use_fabric,
|
||||||
_check_drv(
|
peer_rank=rank,
|
||||||
drv.cuMemSetAccess(
|
|
||||||
int(span_va) + rel,
|
|
||||||
int(alloc_size),
|
|
||||||
[access],
|
|
||||||
1,
|
|
||||||
),
|
|
||||||
"cuMemSetAccess(span)",
|
|
||||||
)
|
)
|
||||||
mapped_chunks.append((rel, int(alloc_size)))
|
mapped_chunks.append((rel, int(alloc_size)))
|
||||||
_check_drv(
|
|
||||||
drv.cuMemRelease(imp_h), "cuMemRelease(span)"
|
|
||||||
)
|
|
||||||
new_mappings.append(
|
new_mappings.append(
|
||||||
(int(span_va), span_size, mapped_chunks)
|
(int(span_va), span_size, mapped_chunks)
|
||||||
)
|
)
|
||||||
@@ -454,7 +628,7 @@ class VmmGraphInputManager:
|
|||||||
self.obj.register_peer_mapped_inputs(peer_ptrs)
|
self.obj.register_peer_mapped_inputs(peer_ptrs)
|
||||||
self._peer_mappings.extend(new_mappings)
|
self._peer_mappings.extend(new_mappings)
|
||||||
except Exception:
|
except Exception:
|
||||||
self._release_peer_mappings(new_mappings)
|
release_mappings(new_mappings)
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
for fd in posix_peer_fds.values():
|
for fd in posix_peer_fds.values():
|
||||||
@@ -472,108 +646,9 @@ class VmmGraphInputManager:
|
|||||||
for fd in local_posix_fds:
|
for fd in local_posix_fds:
|
||||||
os.close(fd)
|
os.close(fd)
|
||||||
for h in retained_handles:
|
for h in retained_handles:
|
||||||
_check_drv(drv.cuMemRelease(h), "cuMemRelease(retained)")
|
check_drv(drv.cuMemRelease(h), "cuMemRelease(retained)")
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
if not self._peer_mappings:
|
if not self._peer_mappings:
|
||||||
return
|
return
|
||||||
self._release_peer_mappings(self._peer_mappings)
|
release_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"
|
|
||||||
)
|
|
||||||
@@ -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,))
|
||||||
Reference in New Issue
Block a user