[NPU] fix: reach torch>=2.8 CUDA memory-pool APIs lazily via torch._C (#29100)

Co-authored-by: Alex Nails <alex.nails@radixark.ai>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Tech Cow
2026-08-29 15:55:14 -07:00
committed by GitHub
co-authored by Alex Nails Claude Opus 5
parent 09ecb9aaaa
commit 3a0f1a1344
4 changed files with 94 additions and 9 deletions
+4
View File
@@ -199,6 +199,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_mps_backends, handle_mps_backends,
handle_nccl_pre_warm, handle_nccl_pre_warm,
handle_npu_backends, handle_npu_backends,
handle_symm_mem_device_support,
handle_xpu_backends, handle_xpu_backends,
) )
@@ -207,6 +208,9 @@ def run_resolution_pipeline(server_args: Any) -> None:
handle_npu_backends(server_args) handle_npu_backends(server_args)
handle_mps_backends(server_args) handle_mps_backends(server_args)
handle_xpu_backends(server_args) handle_xpu_backends(server_args)
# Must precede handle_gpu_memory_settings: its symm-mem prealloc default
# keys off enable_symm_mem.
handle_symm_mem_device_support(server_args)
# OOT platform plugins set fields directly (an interface this tree # OOT platform plugins set fields directly (an interface this tree
# does not own); the diff records what they applied. # does not own); the diff records what they applied.
@@ -76,6 +76,20 @@ def handle_nccl_pre_warm(server_args: Any):
declare_resolution(server_args, "_handle_nccl_pre_warm", pre_warm_nccl=False) declare_resolution(server_args, "_handle_nccl_pre_warm", pre_warm_nccl=False)
def handle_symm_mem_device_support(server_args: Any):
cfg = resolving_view(server_args)
# The symm-mem allocator compiles a CUDA plugin and links -lnccl, so off
# CUDA/HIP (e.g. Ascend NPU) it fails deep in a build step rather than here.
if cfg.enable_symm_mem and not (is_cuda() or is_hip()):
logger.warning(
"--enable-symm-mem is not supported on non CUDA/HIP devices "
"(NCCL symmetric memory is unavailable). Disabling symmetric memory."
)
declare_resolution(
server_args, "_handle_symm_mem_device_support", enable_symm_mem=False
)
def handle_xpu_backends(server_args: Any): def handle_xpu_backends(server_args: Any):
cfg = resolving_view(server_args) cfg = resolving_view(server_args)
if cfg.device == "xpu": if cfg.device == "xpu":
@@ -6,12 +6,10 @@ import traceback
from contextlib import nullcontext from contextlib import nullcontext
import torch import torch
from torch.cuda.memory import (
CUDAPluggableAllocator, # The private _cuda_* pool APIs are absent before torch 2.8; the call sites below
_cuda_beginAllocateCurrentThreadToPool, # reach them via torch._C.<name> so torch 2.7 (Ascend NPU) can still import this.
_cuda_endAllocateToPool, from torch.cuda.memory import CUDAPluggableAllocator
_cuda_releasePool,
)
from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -189,6 +187,11 @@ def get_nccl_mem_pool() -> torch.cuda.MemPool:
All groups share the same pool to avoid memory fragmentation. All groups share the same pool to avoid memory fragmentation.
Comm registration is handled at context exit time. Comm registration is handled at context exit time.
""" """
assert after_2_8_0, (
"--enable-symm-mem requires torch>=2.8 "
"(torch._C._cuda_beginAllocateCurrentThreadToPool was added there)."
)
global _allocator, _mem_pool, _cur_device, _register_func global _allocator, _mem_pool, _cur_device, _register_func
if _allocator is None: if _allocator is None:
import torch.utils.cpp_extension import torch.utils.cpp_extension
@@ -279,7 +282,9 @@ class SymmetricMemoryContext:
_cur_device, _graph_pool_id _cur_device, _graph_pool_id
) )
_cuda_beginAllocateCurrentThreadToPool(self._device_index, self._pool_id) torch._C._cuda_beginAllocateCurrentThreadToPool(
self._device_index, self._pool_id
)
global _active_symmetric_memory_context global _active_symmetric_memory_context
_active_symmetric_memory_context = self _active_symmetric_memory_context = self
@@ -287,8 +292,8 @@ class SymmetricMemoryContext:
return self return self
def __exit__(self, exc_type, exc_val, exc_tb): def __exit__(self, exc_type, exc_val, exc_tb):
_cuda_endAllocateToPool(self._device_index, self._pool_id) torch._C._cuda_endAllocateToPool(self._device_index, self._pool_id)
_cuda_releasePool(self._device_index, self._pool_id) torch._C._cuda_releasePool(self._device_index, self._pool_id)
# Register all unregistered segments # Register all unregistered segments
# with the current comm # with the current comm
self._register_segments_for_comm() self._register_segments_for_comm()
@@ -0,0 +1,62 @@
"""Regression test for https://github.com/sgl-project/sglang/issues/28999.
``pynccl_allocator`` must not import private ``torch.cuda.memory`` symbols at
module scope: they are absent before torch 2.8 and abort startup on Ascend NPU.
"""
import ast
import unittest
from pathlib import Path
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
# test/registered/unit/distributed/<this file> -> repo root
REPO_ROOT = Path(__file__).resolve().parents[4]
SOURCE_PATH = (
REPO_ROOT / "python/sglang/srt/distributed/device_communicators/pynccl_allocator.py"
)
def _import_time_nodes(tree: ast.Module):
"""Yield nodes that run at import time, including ``try`` / ``if`` bodies."""
stack = list(tree.body)
while stack:
node = stack.pop()
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
continue
yield node
stack.extend(ast.iter_child_nodes(node))
class TestPyncclAllocatorImportGuard(CustomTestCase):
def test_no_import_time_private_cuda_memory_symbols(self):
self.assertTrue(
SOURCE_PATH.is_file(),
f"cannot locate pynccl_allocator.py at {SOURCE_PATH}; "
"update REPO_ROOT if the tree layout changed",
)
tree = ast.parse(SOURCE_PATH.read_text(), filename=str(SOURCE_PATH))
offenders = [
alias.name
for node in _import_time_nodes(tree)
if isinstance(node, ast.ImportFrom) and node.module == "torch.cuda.memory"
for alias in node.names
if alias.name.startswith("_cuda_")
]
self.assertEqual(
offenders,
[],
"pynccl_allocator must not import private torch.cuda.memory symbols "
f"at module scope (found {offenders}); these are absent on torch<2.8 "
"and break startup on Ascend NPU. Reach them via torch._C.<name> at "
"the call site instead.",
)
if __name__ == "__main__":
unittest.main()