[ROCm] Use unreg path for aiter custom all-reduce during CUDA graph capture (#20155)

This commit is contained in:
Yuzhen Zhou
2026-03-09 01:09:04 -07:00
committed by GitHub
parent cabe171b6c
commit b719219de9

View File

@@ -4,6 +4,7 @@ import ctypes
import logging
import os
from contextlib import contextmanager
from functools import partial
from typing import Any, List, Optional, Union
import torch
@@ -495,7 +496,11 @@ def dispatch_custom_allreduce():
)
logger.info("[AR] Using AiterCustomAllreduce (AMD default)")
return AiterCustomAllreduce
tms_cudagraph = envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get()
return partial(
AiterCustomAllreduce,
enable_register_for_capturing=not tms_cudagraph,
)
except ImportError as e:
logger.warning(
"[AR] Aiter custom all-reduce not available; "