[AMD] Fix AITER custom all-gather CUDA-graph capture crash under torch_memory_saver (#30557)
Co-authored-by: Xinyu Jiang <xinyuj2@andrew.cmu.edu>
This commit is contained in:
co-authored by
Xinyu Jiang
parent
d74619b373
commit
bd7e54d737
@@ -953,7 +953,8 @@ class GroupCoordinator:
|
|||||||
# 16B alignment, weak-contiguous, supported topology, and per-rank
|
# 16B alignment, weak-contiguous, supported topology, and per-rank
|
||||||
# size <= max_size/(world*2).
|
# size <= max_size/(world*2).
|
||||||
# On a hit, writes directly into the caller's pre-allocated `output` via
|
# On a hit, writes directly into the caller's pre-allocated `output` via
|
||||||
# all_gather_reg during CUDA-graph capture and all_gather_unreg otherwise.
|
# all_gather_reg during CUDA-graph capture, and all_gather_unreg
|
||||||
|
# under torch_memory_saver and other paths.
|
||||||
ca_comm = self.ca_comm
|
ca_comm = self.ca_comm
|
||||||
if (
|
if (
|
||||||
is_hip()
|
is_hip()
|
||||||
@@ -966,6 +967,9 @@ class GroupCoordinator:
|
|||||||
):
|
):
|
||||||
if getattr(ca_comm, "_IS_CAPTURING", False):
|
if getattr(ca_comm, "_IS_CAPTURING", False):
|
||||||
if torch.cuda.is_current_stream_capturing():
|
if torch.cuda.is_current_stream_capturing():
|
||||||
|
if envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get():
|
||||||
|
ca_comm.all_gather_unreg(input, out=output, dim=0)
|
||||||
|
else:
|
||||||
ca_comm.all_gather_reg(input, out=output, dim=0)
|
ca_comm.all_gather_reg(input, out=output, dim=0)
|
||||||
elif is_in_tc_piecewise_cuda_graph():
|
elif is_in_tc_piecewise_cuda_graph():
|
||||||
ca_comm.all_gather_unreg(input, out=output, dim=0)
|
ca_comm.all_gather_unreg(input, out=output, dim=0)
|
||||||
|
|||||||
Reference in New Issue
Block a user