From 379357beaaf45f8554a38ffc7067adafd594547f Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 8 Oct 2026 11:28:34 -0700 Subject: [PATCH 1/3] Distributed tests set device in NCCL init and manually shuts down NCCL Signed-off-by: Tim Moon --- tests/pytorch/distributed/test_fusible_ops.py | 13 +++++++++++++ .../distributed/test_parallel_cross_entropy.py | 1 + 2 files changed, 14 insertions(+) diff --git a/tests/pytorch/distributed/test_fusible_ops.py b/tests/pytorch/distributed/test_fusible_ops.py index c314bf5bb69..ecac9afb876 100644 --- a/tests/pytorch/distributed/test_fusible_ops.py +++ b/tests/pytorch/distributed/test_fusible_ops.py @@ -63,10 +63,20 @@ def world_group() -> torch.distributed.ProcessGroup: init_method=f"file://{os.environ['NVTE_TEST_RDZV_PATH']}", world_size=world_size, rank=rank, + device_id=torch.device("cuda", rank), ) return group +def destroy_world_group() -> None: + """Destroy NCCL process group""" + process_group = world_group() + torch.distributed.barrier(process_group) + torch.cuda.synchronize() + torch.distributed.destroy_process_group(process_group) + world_group.cache_clear() + + def reset_rng(seed: int = 1234) -> None: """Reset random number generators""" torch.manual_seed(seed) @@ -1033,6 +1043,9 @@ def run_parallel_tests() -> None: print(f"Running _test_fp8_scale_update") _test_fp8_scale_update() + # Make sure NCCL shuts down cleanly + destroy_world_group() + # Parallel job sizes _world_sizes = [torch.cuda.device_count()] diff --git a/tests/pytorch/distributed/test_parallel_cross_entropy.py b/tests/pytorch/distributed/test_parallel_cross_entropy.py index 16323154939..a6da464ac77 100644 --- a/tests/pytorch/distributed/test_parallel_cross_entropy.py +++ b/tests/pytorch/distributed/test_parallel_cross_entropy.py @@ -29,6 +29,7 @@ def _run_tensor_parallel(rank, world_size, init_file, label_smoothing): init_method=f"file://{init_file}", rank=rank, world_size=world_size, + device_id=device, ) try: generator = torch.Generator().manual_seed(2025) From 04c91945cb6dbb078754e251b5c4659510b20ebb Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 8 Oct 2026 11:50:57 -0700 Subject: [PATCH 2/3] Fully destroy NCCL environment when finishing distributed TE ops test Signed-off-by: Tim Moon --- tests/pytorch/distributed/test_fusible_ops.py | 23 ++++++++----------- 1 file changed, 10 insertions(+), 13 deletions(-) diff --git a/tests/pytorch/distributed/test_fusible_ops.py b/tests/pytorch/distributed/test_fusible_ops.py index ecac9afb876..0a14d5e352e 100644 --- a/tests/pytorch/distributed/test_fusible_ops.py +++ b/tests/pytorch/distributed/test_fusible_ops.py @@ -68,15 +68,6 @@ def world_group() -> torch.distributed.ProcessGroup: return group -def destroy_world_group() -> None: - """Destroy NCCL process group""" - process_group = world_group() - torch.distributed.barrier(process_group) - torch.cuda.synchronize() - torch.distributed.destroy_process_group(process_group) - world_group.cache_clear() - - def reset_rng(seed: int = 1234) -> None: """Reset random number generators""" torch.manual_seed(seed) @@ -1043,9 +1034,6 @@ def run_parallel_tests() -> None: print(f"Running _test_fp8_scale_update") _test_fp8_scale_update() - # Make sure NCCL shuts down cleanly - destroy_world_group() - # Parallel job sizes _world_sizes = [torch.cuda.device_count()] @@ -1079,7 +1067,16 @@ def main() -> None: parser.add_argument("--parallel", action="store_true", help="Run parallel tests") args = parser.parse_args() if args.parallel: - run_parallel_tests() + try: + run_parallel_tests() + finally: + # Synchronize GPUs + torch.distributed.barrier(world_group()) + torch.cuda.synchronize() + + # Tear down NCCL + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() if __name__ == "__main__": From 19107606c51974186af20b67267621a9fddb5d23 Mon Sep 17 00:00:00 2001 From: Tim Moon <4406448+timmoon10@users.noreply.github.com> Date: Thu, 8 Oct 2026 19:21:28 -0700 Subject: [PATCH 3/3] Apply suggestion from @timmoon10 Signed-off-by: Tim Moon <4406448+timmoon10@users.noreply.github.com> --- tests/pytorch/distributed/test_fusible_ops.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/tests/pytorch/distributed/test_fusible_ops.py b/tests/pytorch/distributed/test_fusible_ops.py index 8cbff804a41..e01d71e45af 100644 --- a/tests/pytorch/distributed/test_fusible_ops.py +++ b/tests/pytorch/distributed/test_fusible_ops.py @@ -1130,12 +1130,9 @@ def main() -> None: run_compile_parallel_tests() else: run_parallel_tests() - finally: - # Synchronize GPUs torch.distributed.barrier(world_group()) torch.cuda.synchronize() - - # Tear down NCCL + finally: if torch.distributed.is_initialized(): torch.distributed.destroy_process_group()