[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:
co-authored by
Alex Nails
Claude Opus 5
parent
09ecb9aaaa
commit
3a0f1a1344
@@ -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()
|
||||||
Reference in New Issue
Block a user