diff --git a/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py b/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py new file mode 100644 index 00000000000..2b3e09d9ad6 --- /dev/null +++ b/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py @@ -0,0 +1,55 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Helpers for parsing ``torch.profiler`` events in the mfsdp_v2 tests.""" + +from torch.autograd import DeviceType +from torch.autograd.profiler_util import FunctionEvent +from torch.profiler import profile as TorchProfiler + + +def events_overlap(first: FunctionEvent, second: FunctionEvent) -> bool: + return ( + first.time_range.start < second.time_range.end + and second.time_range.start < first.time_range.end + ) + + +def collect_linked_kernels( + prof: TorchProfiler, cpu_event_name_substring: str +) -> list[FunctionEvent]: + """Collect device kernel events linked to matching CPU op instances. + + Device events are attributed by their launching CPU op rather than searched by their + own name: device-side names vary across GPU architectures and kernel libraries -- for + example a matmul kernel is named ``nvjet_``/``cutlass_``/``cublas_``... while its + CPU op is simply ``aten::mm``. + + Zero-CTA all-gather copy-engine memcpys are not kernels and are intentionally not + returned. + """ + # A correlation id is shared by a device event and the leaf runtime op that issued it, + # not the enclosing matched op, so walk cpu_parent up from each correlated leaf. Id 0 + # is the "no device correlation" sentinel and is skipped. + events = prof.events() + matching_correlations: set[int] = set() + for event in events: + if event.device_type != DeviceType.CPU or not event.linked_correlation_id: + continue + node = event + while node is not None: + if cpu_event_name_substring in node.name: + matching_correlations.add(event.linked_correlation_id) + break + node = node.cpu_parent + + linked_kernels: list[FunctionEvent] = [] + for event in events: + if event.device_type != DeviceType.CUDA: + continue + if event.activity_type != "kernel": + continue + if event.linked_correlation_id not in matching_correlations: + continue + linked_kernels.append(event) + + return linked_kernels diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py b/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py index 5485a1df954..2d3cc2a677a 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py @@ -8,7 +8,7 @@ import torch import torch.distributed as dist from torch import nn -from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh from torch.distributed.tensor import DTensor from torch.profiler import ProfilerActivity, profile @@ -22,6 +22,10 @@ microbatch, ) from megatron.core.distributed.fsdp.src.megatron_fsdp.mixed_precision import MixedPrecisionPolicy +from tests.unit_tests.distributed.mfsdp_v2.profiler_utils import ( + collect_linked_kernels, + events_overlap, +) logger = logging.getLogger(__name__) @@ -120,11 +124,11 @@ def _mb(num_bytes: int) -> str: return f"{num_bytes / 1024**2:.2f} MB" -def _events_overlap(first, second) -> bool: - return ( - first.time_range.start < second.time_range.end - and second.time_range.start < first.time_range.end - ) +# CPU ops that a device event chains up to via cpu_parent, used to attribute the device +# work to its enclosing collective or matmul operation. +_ALL_GATHER_OP_NAME_SUBSTRING = "allgather" +_REDUCE_SCATTER_OP_NAME_SUBSTRING = "reduce_scatter" +_GEMM_OP_NAME_SUBSTRING = "aten::mm" def _nccl_events(cuda_events, *name_fragments): @@ -479,14 +483,14 @@ def test_root_backward_returns_to_resting_memory(distributed_setup): ) -def test_overlaps_communication_and_compute(distributed_setup): +@pytest.mark.parametrize("use_symm_mem", [False, True], ids=["default", "symmetric_memory"]) +def test_overlaps_communication_and_compute(distributed_setup, use_symm_mem): """Forward and backward communication should overlap GEMM compute.""" world_size = distributed_setup.world_size device = distributed_setup.device if world_size < 2: pytest.skip("This test requires at least 2 ranks.") - mesh = init_device_mesh(device.type, (world_size,)) # A large hidden size keeps the per-layer GEMMs long enough that the # collectives reliably overlap them. The overlap count is otherwise # launch-bound: the host issues kernels with gaps (amplified by CI's @@ -497,12 +501,51 @@ def test_overlaps_communication_and_compute(distributed_setup): dim = 16384 num_children = 4 dtype = torch.bfloat16 + + # new_group requires a default process group. Initialize it here so this test works + # in isolation. Do not eagerly initialize it with device_id in the shared fixture: + # that can hang teardown after communicator splits; see + # https://github.com/pytorch/pytorch/issues/190396. + if not dist.is_initialized(): + dist.init_process_group(backend="nccl") + + if use_symm_mem: + # Dedicated communicator with NCCL's zero-CTA policy. cta_policy is a + # per-communicator property, so scoping it to this group leaves the rest of the + # bucket on default-CTA symmetric-memory kernels (test_symmetric_memory.py asserts + # ncclSymk all-gather kernel counts, which zero-CTA would turn into copy-engine + # memcpys). This 1-D group models the DP (FSDP) sub-mesh that mfsdp is handed in + # production: with EP/TP the full device mesh is multi-dimensional, but mfsdp + # requires an all-FSDP mesh (see experimental/module.py) and never sees the TP/EP + # axes, so only the DP communicator needs the zero-CTA policy. + zero_cta_options = dist.ProcessGroupNCCL.Options() + zero_cta_options.config.cta_policy = dist.ProcessGroupNCCL.NCCL_CTA_POLICY_ZERO + dp_group = dist.new_group(backend="nccl", pg_options=zero_cta_options) + # NCCL window registration can fail when symmetric-memory rendezvous is the first + # operation on a communicator, so initialize this communicator explicitly. + dist.barrier(group=dp_group, device_ids=[device.index]) + else: + dp_group = dist.new_group(backend="nccl") + + mesh = DeviceMesh.from_group(dp_group, device.type) model = MultiChildModel(dim=dim, num_children=num_children).to(dtype=dtype) placements = _flat_placements() policy = MixedPrecisionPolicy(main_params_dtype=dtype, main_grads_dtype=dtype) for layer in model.layers: - fully_shard(layer, mesh=mesh, placements=placements, mixed_precision_policy=policy) - fully_shard(model, mesh=mesh, placements=placements, mixed_precision_policy=policy) + fully_shard( + layer, + mesh=mesh, + placements=placements, + mixed_precision_policy=policy, + use_symm_mem=use_symm_mem, + ) + fully_shard( + model, + mesh=mesh, + placements=placements, + mixed_precision_policy=policy, + use_symm_mem=use_symm_mem, + ) x = torch.randn(4096, dim, device=device, dtype=dtype, requires_grad=True) @@ -521,59 +564,63 @@ def train_one_iteration() -> None: # drop the CUDA events. torch.cuda.synchronize(device) - cuda_events = [event for event in prof.events() if event.device_type.name == "CUDA"] - all_gather_events = _nccl_events(cuda_events, "allgather") - reduce_scatter_events = _nccl_events(cuda_events, "reducescatter", "reduce_scatter") - # GEMM device-kernel names vary across CUDA/cuBLAS versions and GPU archs - # (e.g. "*gemm*", "cutlass*", "cublas*", and cuBLASLt's Hopper "nvjet_sm90_*"). - gemm_events = [ - event - for event in cuda_events - if any(token in event.name.lower() for token in ("gemm", "cutlass", "cublas", "nvjet")) - ] - # Each of the num_children children plus the root all-gathers in forward and - # again in backward, and each reduce-scatters once in backward. - assert len(all_gather_events) == 2 * (num_children + 1), ( - f"Expected {2 * (num_children + 1)} all-gather kernels, " - f"got {[event.name for event in all_gather_events]}." + gemm_kernels = collect_linked_kernels(prof, _GEMM_OP_NAME_SUBSTRING) + # Each child Linear runs one forward and two backward matmuls. aten::mm may also + # launch auxiliary kernels, so check only the matmul lower bound. + assert len(gemm_kernels) >= 3 * num_children, ( + f"Expected at least {3 * num_children} kernels linked to GEMMs, got " + f"{len(gemm_kernels)}: " + f"{[event.name for event in gemm_kernels]}" + ) + + allgather_kernels = collect_linked_kernels(prof, _ALL_GATHER_OP_NAME_SUBSTRING) + reduce_scatter_kernels = collect_linked_kernels(prof, _REDUCE_SCATTER_OP_NAME_SUBSTRING) + # The num_children child layers plus the root are each a sharded module; each does a + # forward and a backward all-gather and one reduce-scatter. Zero-CTA moves the + # all-gather to copy-engine memcpys, so it should not emit all-gather kernels. + num_sharded_modules = num_children + 1 + expected_allgather_kernel_count = 0 if use_symm_mem else 2 * num_sharded_modules + assert len(allgather_kernels) == expected_allgather_kernel_count, ( + f"Expected {expected_allgather_kernel_count} all-gather kernels, got " + f"{len(allgather_kernels)}: {[event.name for event in allgather_kernels]}" ) - assert len(reduce_scatter_events) == num_children + 1, ( - f"Expected {num_children + 1} reduce-scatter kernels, " - f"got {[event.name for event in reduce_scatter_events]}." + assert len(reduce_scatter_kernels) == num_sharded_modules, ( + f"Expected {num_sharded_modules} reduce-scatter kernels, got " + f"{len(reduce_scatter_kernels)}: {[event.name for event in reduce_scatter_kernels]}" ) - assert gemm_events, [event.name for event in cuda_events] - all_gather_streams = {event.device_resource_id for event in all_gather_events} - reduce_scatter_streams = {event.device_resource_id for event in reduce_scatter_events} - gemm_streams = {event.device_resource_id for event in gemm_events} - assert len(all_gather_streams) == 1 + allgather_streams = {event.device_resource_id for event in allgather_kernels} + reduce_scatter_streams = {event.device_resource_id for event in reduce_scatter_kernels} + gemm_streams = {event.device_resource_id for event in gemm_kernels} + if allgather_kernels: + assert len(allgather_streams) == 1 assert len(reduce_scatter_streams) == 1 - assert all_gather_streams.isdisjoint(reduce_scatter_streams) - assert all_gather_streams.isdisjoint(gemm_streams) + assert allgather_streams.isdisjoint(reduce_scatter_streams) + assert allgather_streams.isdisjoint(gemm_streams) assert reduce_scatter_streams.isdisjoint(gemm_streams) - all_gather_overlap_count = sum( - any(_events_overlap(all_gather_event, gemm_event) for gemm_event in gemm_events) - for all_gather_event in all_gather_events + allgather_overlap_count = sum( + any(events_overlap(event, gemm) for gemm in gemm_kernels) for event in allgather_kernels ) reduce_scatter_overlap_count = sum( - any(_events_overlap(reduce_scatter_event, gemm_event) for gemm_event in gemm_events) - for reduce_scatter_event in reduce_scatter_events + any(events_overlap(event, gemm) for gemm in gemm_kernels) + for event in reduce_scatter_kernels ) - # With dim large enough for the GEMMs to dominate launch jitter (see above), - # the prefetched collectives overlap compute deterministically, so assert the - # theoretical maxima (2*(num_children - 1) all-gathers across forward and - # backward, num_children - 1 reduce-scatters in backward) rather than a loose - # floor. - assert all_gather_overlap_count >= 2 * (num_children - 1), ( - f"Expected all-gather to overlap compute, " - f"got {all_gather_overlap_count}/{len(all_gather_events)}." - ) - assert reduce_scatter_overlap_count >= num_children - 1, ( - f"Expected reduce-scatter to overlap compute, " - f"got {reduce_scatter_overlap_count}/{len(reduce_scatter_events)}." + expected_allgather_overlap = 2 * (num_children - 1) + expected_reduce_scatter_overlap = num_children - 1 + if not use_symm_mem: + assert allgather_overlap_count >= expected_allgather_overlap, ( + f"Expected at least {expected_allgather_overlap} all-gathers to " + f"overlap compute, got {allgather_overlap_count}/{len(allgather_kernels)}." + ) + assert reduce_scatter_overlap_count >= expected_reduce_scatter_overlap, ( + f"Expected at least {expected_reduce_scatter_overlap} reduce-scatters to overlap " + f"compute, got {reduce_scatter_overlap_count}/{len(reduce_scatter_kernels)}." ) + # Release the dedicated communicator so it does not leak into the shared session. + dist.destroy_process_group(dp_group) + def test_parameterless_parent_with_child_modules_trains(distributed_setup): """A parent with no unowned parameters should still root trainable child FsdpModules."""