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:
@@ -93,12 +93,14 @@ class CustomAllReduceV2:
|
||||
|
||||
@contextmanager
|
||||
def capture(self):
|
||||
if self.disabled:
|
||||
yield
|
||||
return
|
||||
try:
|
||||
self.obj.set_cuda_graph_capture(True)
|
||||
yield
|
||||
finally:
|
||||
self.obj.set_cuda_graph_capture(False)
|
||||
if not self.disabled:
|
||||
# cannot call when graph is capturing
|
||||
assert (
|
||||
torch.cuda.is_current_stream_capturing() == False
|
||||
|
||||
Reference in New Issue
Block a user