From c8bc23522fe2534b0648f9ce36b7837b38a68f55 Mon Sep 17 00:00:00 2001 From: sky Date: Wed, 6 May 2026 17:07:58 +0800 Subject: [PATCH] Refactor: decouple segment tracking from comm registration (#21392) Signed-off-by: wangfakang --- .../bench_segment_tracking.py | 210 ++++++++++++++++++ .../device_communicators/pynccl_allocator.py | 161 ++++++++++++-- 2 files changed, 348 insertions(+), 23 deletions(-) create mode 100644 benchmark/bench_pynccl_allocator/bench_segment_tracking.py diff --git a/benchmark/bench_pynccl_allocator/bench_segment_tracking.py b/benchmark/bench_pynccl_allocator/bench_segment_tracking.py new file mode 100644 index 000000000..1900fc7f9 --- /dev/null +++ b/benchmark/bench_pynccl_allocator/bench_segment_tracking.py @@ -0,0 +1,210 @@ +""" +Benchmark for comparing CPU overhead of segment tracking methods: +1. nccl_allocator_register_segments_with_comm() - C++ registration with index tracking +2. torch.cuda.memory.memory_snapshot() - PyTorch memory snapshot + +Usage: + python benchmark/bench_pynccl_allocator/bench_segment_tracking.py --num-segments 50 --num-iters 1000 +""" + +import argparse +import time +import warnings +from typing import List + +import torch + +warnings.filterwarnings("ignore") + + +def setup_segments(num_segments: int, segment_size: int = 1024 * 1024): + """ + Allocate a specified number of segments using the NCCL allocator. + """ + import os + + import torch.distributed as dist + + from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + get_nccl_mem_pool, + ) + + # Initialize distributed if not already done + if not dist.is_initialized(): + os.environ.setdefault("MASTER_ADDR", "localhost") + os.environ.setdefault("MASTER_PORT", "29500") + dist.init_process_group( + backend="nccl", + rank=0, + world_size=1, + device_id=torch.device(f"cuda:{torch.cuda.current_device()}"), + ) + + mem_pool = get_nccl_mem_pool() + + # Allocate segments in the pool + tensors: List[torch.Tensor] = [] + with torch.cuda.use_mem_pool(mem_pool): + for _ in range(num_segments): + t = torch.empty(segment_size, dtype=torch.uint8, device="cuda") + tensors.append(t) + + # Keep tensors alive by returning them (caller should hold reference) + return tensors, mem_pool + + +def bench_register_segments_with_comm( + nccl_lib, comm_ptr: int, num_iters: int = 10000 +) -> float: + """ + Benchmark nccl_allocator_register_segments_with_comm() function. + + Args: + nccl_lib: The loaded NCCL allocator library + comm_ptr: The communicator pointer value + num_iters: Number of iterations + + Returns: + Average time per call in microseconds. + """ + import ctypes + + # Setup the C function signature + register_func = nccl_lib.nccl_allocator_register_segments_with_comm + register_func.restype = ctypes.c_int + register_func.argtypes = [ctypes.c_uint64] + + # Warmup + for _ in range(100): + register_func(comm_ptr) + + # Benchmark + start = time.perf_counter() + for _ in range(num_iters): + register_func(comm_ptr) + end = time.perf_counter() + + avg_us = (end - start) / num_iters * 1e6 + return avg_us + + +def bench_mempool_snapshot( + mem_pool: torch.cuda.MemPool, num_iters: int = 10000 +) -> float: + """ + Benchmark torch.cuda.MemPool.snapshot() function. + + Returns: + Average time per call in microseconds. + """ + # Warmup + for _ in range(100): + mem_pool.snapshot() + + # Benchmark + start = time.perf_counter() + for _ in range(num_iters): + mem_pool.snapshot() + end = time.perf_counter() + + avg_us = (end - start) / num_iters * 1e6 + return avg_us + + +def bench_with_various_segment_counts( + segment_counts: List[int], + num_iters: int = 10000, + segment_size: int = 1024 * 1024, # 1MB per segment +): + """ + Run benchmarks with various numbers of tracked segments. + """ + print("=" * 80) + print("Benchmark: Segment Registration CPU Overhead") + print("=" * 80) + print(f"Segment size: {segment_size / 1024 / 1024:.2f} MB") + print(f"Iterations per measurement: {num_iters}") + print() + print( + f"{'Segments':<12} {'register_segments (µs)':<30} {'snapshot (µs)':<20} {'Speedup':<10}" + ) + print("-" * 80) + + all_tensors = [] # Keep all tensors alive + comm_ptr = 0 # Use dummy comm_ptr for benchmarking (no actual NCCL registration) + + for num_segments in segment_counts: + # Clean up previous segments + all_tensors = [] + + # Allocate segments (this initializes _nccl_allocator_lib via get_nccl_mem_pool) + tensors, mem_pool = setup_segments(num_segments, segment_size) + all_tensors.extend(tensors) + + # Sync to ensure allocations are complete + torch.cuda.synchronize() + + # Import _nccl_allocator_lib after setup_segments (ensures library is loaded) + from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + _nccl_allocator_lib, + ) + + # Run benchmarks + time_register = bench_register_segments_with_comm( + _nccl_allocator_lib, comm_ptr, num_iters + ) + time_snapshot = bench_mempool_snapshot(mem_pool, num_iters) + + speedup = time_snapshot / time_register if time_register > 0 else float("inf") + + print( + f"{num_segments:<12} {time_register:<30.3f} {time_snapshot:<20.3f} {speedup:<10.2f}x" + ) + + print("-" * 80) + print() + + +def main(): + parser = argparse.ArgumentParser( + description="Benchmark segment tracking methods in pynccl_allocator" + ) + parser.add_argument( + "--num-segments", + type=int, + nargs="+", + default=[10, 50, 100, 200, 500, 1000], + help="Number of segments to track (can specify multiple values)", + ) + parser.add_argument( + "--num-iters", + type=int, + default=10000, + help="Number of iterations for each measurement", + ) + parser.add_argument( + "--segment-size", + type=int, + default=1024 * 1024, # 1MB + help="Size of each segment in bytes", + ) + args = parser.parse_args() + + # Check CUDA availability + if not torch.cuda.is_available(): + print("Error: CUDA is not available. This benchmark requires a GPU.") + return + + # Initialize CUDA context by creating a small tensor + _ = torch.zeros(1, device="cuda") + + # Run benchmarks + bench_with_various_segment_counts( + segment_counts=args.num_segments, + num_iters=args.num_iters, + segment_size=args.segment_size, + ) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py index 55b14ccbc..efab33daf 100644 --- a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py +++ b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py @@ -1,3 +1,4 @@ +import ctypes import logging import os import tempfile @@ -19,9 +20,23 @@ from sglang.srt.utils.common import torch_release after_2_8_0 = torch_release >= (2, 8) +# C++ source for the NCCL allocator plugin +# Key design: +# 1. nccl_alloc_plug: Allocates memory via ncclMemAlloc and TRACKS the segment +# (ptr, size). Does NOT register with any comm at allocation time. +# 2. nccl_free_plug: Frees memory via ncclMemFree and UNTRACKS the segment. +# Each segment is tracked only during its lifetime (from alloc to free). +# 3. Segment tracking uses thread-safe std::vector + unordered_map for O(1) operations. +# 4. Registration via nccl_allocator_register_segments_with_comm: Registers all +# tracked segments with a given comm, using index-based tracking to avoid +# re-registration. Registration state is maintained per-communicator in C++. nccl_allocator_source = """ #include +#include +#include +#include +#include extern "C" { @@ -35,13 +50,16 @@ typedef enum { ncclSuccess = 0, ncclRemoteError = 6, ncclInProgress = 7, ncclNumResults = 8 } ncclResult_t; + +// NCCL symmetric memory window flags +#define NCCL_WIN_COLL_SYMMETRIC 0x01 + typedef struct ncclComm* ncclComm_t; typedef struct ncclWindow_vidmem* ncclWindow_t; -ncclResult_t ncclCommWindowRegister(ncclComm_t comm, void* buff, size_t size, ncclWindow_t* win, int winFlags); -#define NCCL_WIN_COLL_SYMMETRIC 0x01 ncclResult_t ncclMemAlloc(void** ptr, size_t size); ncclResult_t ncclMemFree(void *ptr); +ncclResult_t ncclCommWindowRegister(ncclComm_t comm, void* buff, size_t size, ncclWindow_t* win, int winFlags); const char* ncclGetErrorString(ncclResult_t result); #define NCCLCHECK(cmd) do { \ @@ -53,23 +71,77 @@ const char* ncclGetErrorString(ncclResult_t result); } \ } while(0) +// Segment information structure +struct Segment { + void* ptr; + size_t size; + Segment(void* p, size_t s) : ptr(p), size(s) {} +}; + +// Thread-safe segment tracking +// Segment tracking using std::vector for FIFO order. +// g_segments is maintained in insertion order (oldest first). +static std::vector g_segments; +static std::mutex g_segment_mutex; + +// Track which segments have been registered with each communicator. +// Key: comm_ptr, Value: the next segment index to register for this comm. +static std::unordered_map g_comm_registration_index; + +// Add a segment to the tracking (appends to end, maintaining FIFO order) +static void track_segment(void* ptr, size_t size) { + std::lock_guard lock(g_segment_mutex); + g_segments.emplace_back(ptr, size); +} + void* nccl_alloc_plug(size_t size, int device, void* stream) { - void* ptr; - NCCLCHECK(ncclMemAlloc(&ptr, size)); + void* ptr; + NCCLCHECK(ncclMemAlloc(&ptr, size)); - const char *str_val = getenv("SGLANG_TMP_NCCL_COMM_VALUE"); - char *endptr; - void* int_val = (void *)strtoull(str_val, &endptr, 0); + // Track the segment but do NOT register with any comm + // Registration will be done at context exit via register_segments_with_comm + track_segment(ptr, size); - ncclComm_t comm = (ncclComm_t)(int_val); - ncclWindow_t win; - NCCLCHECK(ncclCommWindowRegister(comm, ptr, size, &win, NCCL_WIN_COLL_SYMMETRIC)); - - return ptr; + return ptr; } void nccl_free_plug(void* ptr, size_t size, int device, void* stream) { - ncclResult_t err = ncclMemFree(ptr); + ncclResult_t err = ncclMemFree(ptr); + // NOTE: We assume that no individual allocation will be freed until the + // entire memory pool is destroyed. If this assumption does not hold, + // we will encounter asymmetry issues between GPUs. For now, we clear + // all tracking state when the pool is destroyed. + std::lock_guard lock(g_segment_mutex); + g_segments = std::vector(); + g_comm_registration_index = std::unordered_map(); +} + +// Register all tracked segments with a communicator. +// Uses an index-based approach to avoid re-registering already-registered segments. +// Returns 0 on success, non-zero on failure. +int nccl_allocator_register_segments_with_comm(uintptr_t comm_ptr) { + std::lock_guard lock(g_segment_mutex); + + ncclComm_t comm = reinterpret_cast(comm_ptr); + + // Get the starting index for this communicator + size_t start_index = g_comm_registration_index[comm_ptr]; + + // Register all segments from start_index to the current end + for (size_t i = start_index; i < g_segments.size(); ++i) { + const Segment& seg = g_segments[i]; + ncclWindow_t win; + ncclResult_t res = ncclCommWindowRegister(comm, seg.ptr, seg.size, &win, NCCL_WIN_COLL_SYMMETRIC); + if (res != ncclSuccess) { + fprintf(stderr, "ERROR: NCCL symmetric memory registration failed. '%s'\\n", ncclGetErrorString(res)); + return res; + } + } + + // Update the registration index for this communicator + g_comm_registration_index[comm_ptr] = g_segments.size(); + + return ncclSuccess; } } @@ -81,6 +153,9 @@ _graph_pool_id = None _cur_device = None _active_symmetric_memory_context = None +# Reference to the C registration function (with arg types set) +_register_func = None + def is_symmetric_memory_enabled(): try: @@ -107,9 +182,15 @@ def restore_symmetric_memory_context(saved_context): saved_context.__enter__() -def get_nccl_mem_pool(): - global _allocator, _mem_pool, _cur_device - if _mem_pool is None: +def get_nccl_mem_pool() -> torch.cuda.MemPool: + """ + Get the shared MemPool for all groups. + + All groups share the same pool to avoid memory fragmentation. + Comm registration is handled at context exit time. + """ + global _allocator, _mem_pool, _cur_device, _register_func + if _allocator is None: import torch.utils.cpp_extension out_dir = os.path.join(tempfile.gettempdir(), "symm_allocator") @@ -124,7 +205,7 @@ def get_nccl_mem_pool(): torch.distributed.barrier() nccl_allocator_libname = "nccl_allocator" - torch.utils.cpp_extension.load_inline( + lib_path = torch.utils.cpp_extension.load_inline( name=nccl_allocator_libname, cpp_sources=nccl_allocator_source, with_cuda=True, @@ -133,6 +214,7 @@ def get_nccl_mem_pool(): is_python_module=False, build_directory=out_dir, ) + nccl_allocator_lib = ctypes.CDLL(lib_path) _allocator = CUDAPluggableAllocator( f"{out_dir}/{nccl_allocator_libname}.so", "nccl_alloc_plug", @@ -140,6 +222,12 @@ def get_nccl_mem_pool(): ).allocator() _mem_pool = torch.cuda.MemPool(_allocator) _cur_device = torch.cuda.current_device() + + # Setup the C function for registration with correct arg types + _register_func = nccl_allocator_lib.nccl_allocator_register_segments_with_comm + _register_func.restype = ctypes.c_int + _register_func.argtypes = [ctypes.c_uint64] + return _mem_pool @@ -151,6 +239,14 @@ class SymmetricMemoryContext: by `ncclMemAlloc` and registered by `ncclCommWindowRegister`. Due to this, we introduce this context manager. All tensors created under this context will be correctly allocated and registered with a custom allocator. + + Key design: + - All groups share a single MemPool to avoid memory fragmentation. + - At allocation time, ptrs are tracked but NOT registered with any comm. + - At context exit time, nccl_allocator_register_segments_with_comm is called + to register all tracked segments with the current comm. The C++ layer + tracks per-comm registration state using index-based tracking to avoid + re-registration of already-registered segments. """ def __init__( @@ -162,10 +258,14 @@ class SymmetricMemoryContext: self._device_index = torch.cuda.current_device() self.is_graph_capture = torch.cuda.is_current_stream_capturing() + # Get comm ptr for tracking registrations + # Use the comm pointer value as unique identifier + self._comm_ptr = self.group_coordinator.pynccl_comm.comm.value + def __enter__(self): assert ( self.group_coordinator.pynccl_comm is not None - ), f"Symmetric memory requires pynccl to be enabled in group '{self.group_coordinator.group_name}'" + ), f"Symmetric memory requires pynccl to be enabled in group '{self.group_coordinator.unique_name}'" if self.is_graph_capture: assert ( @@ -181,11 +281,6 @@ class SymmetricMemoryContext: _cuda_beginAllocateCurrentThreadToPool(self._device_index, self._pool_id) - # Set the env var to pass this argument to the C functions. - os.environ["SGLANG_TMP_NCCL_COMM_VALUE"] = str( - self.group_coordinator.pynccl_comm.comm.value - ) - global _active_symmetric_memory_context _active_symmetric_memory_context = self @@ -194,6 +289,9 @@ class SymmetricMemoryContext: def __exit__(self, exc_type, exc_val, exc_tb): _cuda_endAllocateToPool(self._device_index, self._pool_id) _cuda_releasePool(self._device_index, self._pool_id) + # Register all unregistered segments + # with the current comm + self._register_segments_for_comm() if self.is_graph_capture: if after_2_8_0: @@ -206,6 +304,23 @@ class SymmetricMemoryContext: global _active_symmetric_memory_context _active_symmetric_memory_context = None + def _register_segments_for_comm(self): + """ + Register all tracked segments with the current comm. + + Delegates to C++ layer which handles: + 1. Tracking which segments have been registered with each comm + 2. Only registering new segments (avoiding re-registration) + 3. Thread-safe access to the segment registry + """ + + # Call C++ API to register all segments with this comm + # C++ layer tracks per-comm registration state internally + result = _register_func(self._comm_ptr) + assert ( + result == 0 + ), f"nccl_allocator_register_segments_with_comm failed with return code: {result}" + def use_symmetric_memory(group_coordinator: GroupCoordinator, disabled: bool = False): disabled = (