diff --git a/tests/pytorch/distributed/test_fusible_ops.py b/tests/pytorch/distributed/test_fusible_ops.py index 8057a75a84..e01d71e45a 100644 --- a/tests/pytorch/distributed/test_fusible_ops.py +++ b/tests/pytorch/distributed/test_fusible_ops.py @@ -63,6 +63,7 @@ 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 @@ -1124,10 +1125,16 @@ def main() -> None: if args.compile and not args.parallel: parser.error("--compile requires --parallel") if args.parallel: - if args.compile: - run_compile_parallel_tests() - else: - run_parallel_tests() + try: + if args.compile: + run_compile_parallel_tests() + else: + run_parallel_tests() + torch.distributed.barrier(world_group()) + torch.cuda.synchronize() + finally: + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() if __name__ == "__main__": diff --git a/tests/pytorch/distributed/test_parallel_cross_entropy.py b/tests/pytorch/distributed/test_parallel_cross_entropy.py index 1632315493..a6da464ac7 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)