From b719219de9fdd4306f04b24a79c71366c1c0a323 Mon Sep 17 00:00:00 2001 From: Yuzhen Zhou <82826991+zyzshishui@users.noreply.github.com> Date: Mon, 9 Mar 2026 01:09:04 -0700 Subject: [PATCH] [ROCm] Use unreg path for aiter custom all-reduce during CUDA graph capture (#20155) --- .../distributed/device_communicators/custom_all_reduce.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py b/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py index b6ca22ca4..5e852a080 100644 --- a/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py +++ b/python/sglang/srt/distributed/device_communicators/custom_all_reduce.py @@ -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; "