From 2c72323d90e157866e7560baf4dc28efa44aa158 Mon Sep 17 00:00:00 2001 From: Zhiyao Jiang Date: Mon, 10 Aug 2026 18:36:15 -0400 Subject: [PATCH] [AMD] Fix AITER custom reduce-scatter CUDA-graph capture crash under torch_memory_saver (#34203) --- python/sglang/srt/distributed/parallel_state.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index aa25d7ddb..f1d134e91 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -1080,7 +1080,10 @@ class GroupCoordinator: return False if getattr(ca_comm, "_IS_CAPTURING", False): if torch.cuda.is_current_stream_capturing(): - ca_comm.reduce_scatter(input, output, registered=True) + if envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get(): + ca_comm.reduce_scatter(input, output, registered=False) + else: + ca_comm.reduce_scatter(input, output, registered=True) elif is_in_tc_piecewise_cuda_graph(): ca_comm.reduce_scatter(input, output, registered=False) else: