Followup fix for Custom AR V2 in non NVL scenarios (#24742)

Co-authored-by: b8zhong <b8zhong@users.noreply.github.com>
This commit is contained in:
Brayden Zhong
2026-05-10 16:57:46 -07:00
committed by GitHub
co-authored by b8zhong
parent b202778aa5
commit 8acb0270fd
@@ -93,24 +93,26 @@ class CustomAllReduceV2:
@contextmanager @contextmanager
def capture(self): def capture(self):
if self.disabled:
yield
return
try: try:
self.obj.set_cuda_graph_capture(True) self.obj.set_cuda_graph_capture(True)
yield yield
finally: finally:
self.obj.set_cuda_graph_capture(False) self.obj.set_cuda_graph_capture(False)
if not self.disabled: # cannot call when graph is capturing
# cannot call when graph is capturing assert (
assert ( torch.cuda.is_current_stream_capturing() == False
torch.cuda.is_current_stream_capturing() == False ), "Cannot register graph inputs while capturing CUDA graph"
), "Cannot register graph inputs while capturing CUDA graph" pairs = self.obj.share_graph_inputs()
pairs = self.obj.share_graph_inputs() handles = [handle for _, handle in pairs]
handles = [handle for _, handle in pairs] offsets = [offset for offset, _ in pairs]
offsets = [offset for offset, _ in pairs] handles_all = self._share_list(handles)
handles_all = self._share_list(handles) offsets_all = self._share_list(offsets)
offsets_all = self._share_list(offsets) result = [list(zip(o, h)) for o, h in zip(offsets_all, handles_all)]
result = [list(zip(o, h)) for o, h in zip(offsets_all, handles_all)] self.obj.register_inputs(result)
self.obj.register_inputs(result) log_info_on_rank0(logger, f"Registering {len(pairs)} cuda graph addresses")
log_info_on_rank0(logger, f"Registering {len(pairs)} cuda graph addresses")
def should_custom_ar(self, inp: torch.Tensor) -> bool: def should_custom_ar(self, inp: torch.Tensor) -> bool:
"""Check if the input tensor is suitable for custom all-reduce.""" """Check if the input tensor is suitable for custom all-reduce."""