diff --git a/CHANGELOG.md b/CHANGELOG.md index 878f0b76e4..16bb833938 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -35,6 +35,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 cells out of a mesh with hundreds of millions of vertices). Index normalization avoids allocating a full-mesh range and preserves empty slices, integer indices, and boolean masks. Point fields use ordinary indexed gathers. +- `index_select` and integer-tensor indexing on a `ShardTensor` exchange only + the requested rows instead of all-gathering the source; no host sync, and + `torch.compile` safe. - The stream-based ring helpers `perform_ring_iteration_async` and `get_comm_stream` are removed; use `perform_ring_iteration_funcol` with `wait=False` and `finish_ring_iteration` for overlap. @@ -88,6 +91,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - `torch.unbind` on a `Partial` `ShardTensor` resolves the pending reduction before slicing instead of silently dropping it. - Fixed autocasting bugs for some attention operations in domain parallelism. +- Backward through `ShardTensor.to_local()` adopts the gradient's shard order + along with its placements, so a sharded primal with a replicated cotangent + no longer trips a `DTensorSpec` assertion. ### Security diff --git a/physicsnemo/domain_parallel/shard_tensor.py b/physicsnemo/domain_parallel/shard_tensor.py index b61733d23d..56e6fbe8d9 100644 --- a/physicsnemo/domain_parallel/shard_tensor.py +++ b/physicsnemo/domain_parallel/shard_tensor.py @@ -183,10 +183,14 @@ def backward(ctx, grad_output: DTensor): # the gradient's tensor_meta: the grad shares the primal's shape but # not necessarily its stride (grad-of-permute is permute-of-grad), # and stamping the primal's stride onto differently-laid-out grad - # memory breaks downstream .view() calls. + # memory breaks downstream .view() calls. The shard order must follow + # the placements: DTensorSpec asserts that every Shard placement has a + # matching entry, so a Shard primal with a Replicate cotangent (a + # replicated result gathered from sharded rows) would otherwise fail. cached_spec = dataclasses.replace( cached_spec, placements=grad_placements, + shard_order=grad_output._spec.shard_order, tensor_meta=grad_output._spec.tensor_meta, ) return (_dtensor_to_shard_tensor(grad_output, cached_spec),) diff --git a/physicsnemo/domain_parallel/shard_utils/__init__.py b/physicsnemo/domain_parallel/shard_utils/__init__.py index ccb0121a8b..5cb07cd519 100644 --- a/physicsnemo/domain_parallel/shard_utils/__init__.py +++ b/physicsnemo/domain_parallel/shard_utils/__init__.py @@ -33,6 +33,7 @@ def register_shard_wrappers(): from .attention_patches import sdpa_wrapper from .conv_patches import generic_conv_nd_wrapper from .index_ops import ( + getitem_wrapper, index_select_wrapper, sharded_select_backward_helper, sharded_select_helper, diff --git a/physicsnemo/domain_parallel/shard_utils/exchange.py b/physicsnemo/domain_parallel/shard_utils/exchange.py new file mode 100644 index 0000000000..52f301451b --- /dev/null +++ b/physicsnemo/domain_parallel/shard_utils/exchange.py @@ -0,0 +1,193 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023 - 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +r"""Variable-size row exchange primitives shared by the ShardTensor row ops. + +Both the halo scatter correction (:mod:`halo_scatter`) and the routed gather +(:mod:`index_ops`) move rows between ranks with a functional-collective +``all_to_all`` whose split sizes are decided at run time, fold contributions +with a configurable accumulator, and hand a plain group-name token to +``custom_op`` bodies. Those pieces live here; the routing that decides *which* +rows move stays with each op. +""" + +from __future__ import annotations + +import torch +import torch.distributed as dist +import torch.distributed._functional_collectives as funcol +from torch.distributed.device_mesh import DeviceMesh + +__all__ = [ + "funcol_all_to_all_v_rows", + "get_fp32_scatter_accumulator", + "resolve_group_name", + "set_fp32_scatter_accumulator", +] + + +def _funcol_group_arg(group: object) -> object: + r"""Return *group* in the form functional collectives accept. + + Parameters + ---------- + group : DeviceMesh or ProcessGroup or str or None + Group specification. ``None`` means the default world group, which + funcol does not accept directly. + + Returns + ------- + object + ``(mesh, 0)`` for a mesh, the default group for ``None``, otherwise + *group* unchanged. + """ + if isinstance(group, DeviceMesh): + return (group, 0) + if group is None: + return dist.distributed_c10d._get_default_group() + return group + + +# Accumulator used when folding ``float32`` scatter contributions (halo +# corrections, routed scatter-add). ``float64`` by default: row folds sum many +# contributions and downstream applications rely on the extra significance. +# ``float32`` trades that for speed and memory on the accumulation buffer. +_FP32_SCATTER_ACCUMULATOR: torch.dtype = torch.float64 + + +def set_fp32_scatter_accumulator(dtype: torch.dtype) -> None: + r"""Choose the accumulator for ``float32`` scatter folds. + + Parameters + ---------- + dtype : torch.dtype + ``torch.float64`` (default; highest significance) or ``torch.float32`` + (accumulate in place; cheaper on large local shards). + + Raises + ------ + ValueError + For any other dtype. + """ + global _FP32_SCATTER_ACCUMULATOR + if dtype not in (torch.float32, torch.float64): + raise ValueError( + f"fp32 scatter accumulator must be float32 or float64, got {dtype}" + ) + _FP32_SCATTER_ACCUMULATOR = dtype + + +def get_fp32_scatter_accumulator() -> torch.dtype: + r"""Accumulator currently used for ``float32`` scatter folds. + + Returns + ------- + torch.dtype + ``torch.float64`` unless changed by :func:`set_fp32_scatter_accumulator`. + """ + return _FP32_SCATTER_ACCUMULATOR + + +def _accumulator_dtype(dtype: torch.dtype) -> torch.dtype: + r"""Dtype for folding scatter contributions of *dtype*. + + Row folds sum many contributions, so accumulating in the input precision + loses significance for the smaller float types. + + Parameters + ---------- + dtype : torch.dtype + Dtype of the contributions. + + Returns + ------- + torch.dtype + :func:`get_fp32_scatter_accumulator` for ``float32``, ``float32`` for + ``float16`` / ``bfloat16``, otherwise *dtype* itself. + """ + if dtype == torch.float32: + return _FP32_SCATTER_ACCUMULATOR + if dtype in (torch.float16, torch.bfloat16): + return torch.float32 + return dtype + + +def resolve_group_name(group: object) -> str: + r"""Resolve *group* to a c10d group-name string. + + The name is the traceable token a ``custom_op`` can carry in place of a + ``ProcessGroup``. + + Parameters + ---------- + group : DeviceMesh or ProcessGroup or str or None + A 1-D device mesh (its dim-0 group is used), a process group, an + existing group name, or ``None`` for the default world group. + + Returns + ------- + str + The group name; ``""`` denotes the default world group. + """ + if group is None: + return "" + if isinstance(group, str): + return group + if isinstance(group, DeviceMesh): + return group.get_group(0).group_name + return group.group_name + + +def funcol_all_to_all_v_rows( + send_rows: torch.Tensor, + send_counts: list[int], + recv_counts: list[int], + group: object = None, +) -> torch.Tensor: + r"""AOT-traceable variable-size row ``all_to_all``. + + Parameters + ---------- + send_rows : torch.Tensor + Rows to send, shape ``(sum(send_counts), *trailing)``, ordered by + destination rank. + send_counts : list[int] + Number of rows sent to each rank. + recv_counts : list[int] + Number of rows received from each rank. + group : DeviceMesh or ProcessGroup or str or None, optional + Group to exchange over; ``None`` is the default world group. + + Returns + ------- + torch.Tensor + Received rows, shape ``(sum(recv_counts), *trailing)``, ordered by + source rank. + """ + trailing = tuple(send_rows.shape[1:]) + row_size = 1 + for d in trailing: + row_size *= d + flat_send = send_rows.contiguous().reshape(-1) + send_flat = [c * row_size for c in send_counts] + recv_flat = [c * row_size for c in recv_counts] + total_recv = sum(recv_counts) + flat_recv = funcol.wait_tensor( + funcol.all_to_all_single( + flat_send, recv_flat, send_flat, _funcol_group_arg(group) + ) + ) + return flat_recv.reshape((total_recv,) + trailing) diff --git a/physicsnemo/domain_parallel/shard_utils/halo_scatter.py b/physicsnemo/domain_parallel/shard_utils/halo_scatter.py index 2146c36470..54bb55281e 100644 --- a/physicsnemo/domain_parallel/shard_utils/halo_scatter.py +++ b/physicsnemo/domain_parallel/shard_utils/halo_scatter.py @@ -32,9 +32,17 @@ import torch import torch.distributed as dist -import torch.distributed._functional_collectives as funcol from torch.distributed.device_mesh import DeviceMesh +from physicsnemo.domain_parallel.shard_utils.exchange import ( # noqa: F401 + _accumulator_dtype, + _funcol_group_arg, + funcol_all_to_all_v_rows, + get_fp32_scatter_accumulator, + resolve_group_name, + set_fp32_scatter_accumulator, +) + __all__ = [ "funcol_all_to_all_v_rows", "halo_forward_exchange", @@ -46,62 +54,6 @@ ] -def _funcol_group_arg(group: object) -> object: - r"""Return *group* in the form functional collectives accept (funcol rejects ``None``).""" - if isinstance(group, DeviceMesh): - return (group, 0) - if group is None: - return dist.distributed_c10d._get_default_group() - return group - - -def _accumulator_dtype(dtype: torch.dtype) -> torch.dtype: - r"""Higher-precision dtype for folding scatter contributions: ``float32`` accumulates - in ``float64`` and reduced-precision (``float16``/``bfloat16``) in ``float32``; all - other dtypes accumulate in place. Row folds sum many contributions, so accumulating in - the input precision loses significance for the smaller float types.""" - if dtype == torch.float32: - return torch.float64 - if dtype in (torch.float16, torch.bfloat16): - return torch.float32 - return dtype - - -def _halo_group_name(group: object) -> str: - r"""Resolve *group* to a c10d group-name string (``""`` = default world group), the - traceable token a ``custom_op`` can carry in place of a ``ProcessGroup``.""" - if group is None: - return "" - if isinstance(group, str): - return group - if isinstance(group, DeviceMesh): - return group._dim_group_names[0] - return group.group_name - - -def funcol_all_to_all_v_rows( - send_rows: torch.Tensor, - send_counts: list[int], - recv_counts: list[int], - group: object = None, -) -> torch.Tensor: - r"""AOT-traceable variable-sized row ``all_to_all`` via ``funcol.all_to_all_single``; send/recv buffers are destination/source-rank ordered.""" - trailing = tuple(send_rows.shape[1:]) - row_size = 1 - for d in trailing: - row_size *= d - flat_send = send_rows.contiguous().reshape(-1) - send_flat = [c * row_size for c in send_counts] - recv_flat = [c * row_size for c in recv_counts] - total_recv = sum(recv_counts) - flat_recv = funcol.wait_tensor( - funcol.all_to_all_single( - flat_send, recv_flat, send_flat, _funcol_group_arg(group) - ) - ) - return flat_recv.reshape((total_recv,) + trailing) - - # A transport backend owns the whole reverse/forward exchange, since the data-movement # structure -- not just the collective call -- is transport-specific. @@ -194,7 +146,7 @@ def _symm_group_name(group: object) -> str: if isinstance(group, str): return group if isinstance(group, DeviceMesh): - return group._dim_group_names[0] + return group.get_group(0).group_name return group.group_name @@ -613,7 +565,7 @@ def halo_scatter_correct( contributions into their owners and refresh the ghost rows as a single AOT-traceable, differentiable graph node, with ``routing`` (from :func:`pack_halo_routing`) riding as a graph input to survive graph breaks.""" - return _halo_scatter_correct_op(padded, routing, _halo_group_name(group)) + return _halo_scatter_correct_op(padded, routing, resolve_group_name(group)) # ShardTensor scatter/index-add integration. A ShardTensor carrying packed routing as an diff --git a/physicsnemo/domain_parallel/shard_utils/index_ops.py b/physicsnemo/domain_parallel/shard_utils/index_ops.py index 0bf32a72d4..749d560422 100644 --- a/physicsnemo/domain_parallel/shard_utils/index_ops.py +++ b/physicsnemo/domain_parallel/shard_utils/index_ops.py @@ -16,9 +16,12 @@ from __future__ import annotations +import itertools +import math from typing import Any, Callable import torch +from torch.distributed.tensor import DTensor from torch.distributed.tensor.placement_types import ( Replicate, Shard, @@ -30,6 +33,15 @@ TensorMeta, _stride_from_contiguous_shape_C_style, ) +from physicsnemo.domain_parallel.shard_tensor import ( + _is_tracing, + _torch_function_fallback_via_dtensor, +) +from physicsnemo.domain_parallel.shard_utils.exchange import ( + _accumulator_dtype, + funcol_all_to_all_v_rows, + resolve_group_name, +) from physicsnemo.domain_parallel.shard_utils.patch_core import ( MissingShardPatch, ) @@ -37,207 +49,670 @@ aten = torch.ops.aten -class ShardedIndexSelect(torch.autograd.Function): - r"""Autograd function implementing a differentiable index_select operation for ShardTensors. +# --------------------------------------------------------------------------- +# Routed gather / scatter-add along the sharded dim +# --------------------------------------------------------------------------- +# +# ``values`` is ``Shard(d)`` with known per-rank sizes; ``index`` holds GLOBAL +# positions along ``d`` that may live on any rank. Everything about the exchange +# is known from the specs before any data moves: the shard boundaries of +# ``values`` (``offsets``) and how many requests every rank holds +# (``request_sizes``, from the index's sharding shapes). So the exchange uses +# STATIC split sizes and no host sync: +# +# 1. Every rank sends its full request list to every owner (sizes: +# ``request_sizes``). +# 2. Each owner gathers, for each requester, a full-length row buffer in +# request order -- rows it does not own are filler -- and sends it back. +# 3. The requester holds ``W`` buffers of its own length and picks, per +# request, the row from the owner it computed locally (``bucketize``). +# +# Bandwidth is ``W`` x the requested rows instead of 1 x plus a count exchange; +# for the small ``W`` of a domain group that is cheaper than stalling the CPU +# on a device-to-host copy twice per gather. The backward is the mirror image: +# the requester sends its gradient buffer to every owner, each owner +# ``index_add``s only the rows it owns, which yields a complete ``Shard(d)`` +# gradient with no pending reduction. The exchanged requests are saved from +# the forward, so the backward is a single collective. +# +# Transport (row all-to-all with given split sizes), group-name resolution and +# the accumulation policy come from ``shard_utils.exchange`` (shared with +# ``halo_scatter``); only the routing is specific to this module. +# +# The routed core is a ``custom_op`` (opaque to fake mode), so the op is +# ``torch.compile`` safe: every input size is a spec constant and every output +# shape is static (index shape + trailing dims, or the local values shape for +# the backward). + +_INDEX_DTYPES = (torch.int32, torch.int64) + + +def _wrap_negative(index_flat: torch.Tensor, n_global: int) -> torch.Tensor: + r"""Wrap negative global positions like eager indexing does. + + Parameters + ---------- + index_flat : torch.Tensor + 1-D global positions; negatives count from the end. + n_global : int + Global extent along the indexed dim. Positions ``>= n_global`` are not + checked (they select filler rows on the owner side). - This class provides both forward and backward pass implementations to enable - gradient computation through the index_select operation when working with - distributed sharded tensors. + Returns + ------- + torch.Tensor + Non-negative positions. """ + return torch.where(index_flat < 0, index_flat + n_global, index_flat) - @staticmethod - def forward( - tensor: ShardTensor, - dim: int, - index: ShardTensor, - ) -> ShardTensor: - r"""Implement a differentiable index select operation on ShardTensors. - - This requires collectives and temporarily utilizing the full shape. - It could be optimized, for large tensors, to use a ring and smarter indexing. - - Parameters - ---------- - tensor : ShardTensor - Input tensor to select from. - dim : int - Dimension along which to index. - index : ShardTensor - Indices to select. - - Returns - ------- - ShardTensor - Output tensor containing the selected elements. - - Raises - ------ - MissingShardPatch - If the index sharding strategy is not implemented. - """ - # This is the simplest implementation, to enable functionality. - # It could be optimized for very large tensors to ensure performace. - - # First - Make sure we have the full input tensor - # Triggers an all_gather(_v) for (uneven) tensors. - local_tensor = tensor.full_tensor() - - # Perform the index select using the local values of the index: - local_index = index.to_local() - - # Get everything requested from the local index: - local_values = aten.index_select(local_tensor, dim, local_index) - - # Now, we do gymnastics to make sure the output is correctly sharded. - # Because index is one dimensional, by requirement of the underlying function, - # it's not as annoying as it could be. - index_placement = index._spec.placements[0] - - if index_placement.is_shard(): - # Then, we return a tensor sharded along dim aka Shard(dim). - # Size per rank is easy to compute, no communication needed. - output_size = list(tensor.shape) - output_shard_sizes = {} - for mesh_dim, index_shard_sizes in index._spec.sharding_shapes().items(): - output_shard_sizes[mesh_dim] = [] - for local_chunk_size in index_shard_sizes: - this_shard_size = output_size - this_shard_size[dim] = local_chunk_size[0] - # Plain int tuples (never torch.Size) -- see - # ShardTensorSpec._sharding_shapes field docs. - output_shard_sizes[mesh_dim].append(tuple(this_shard_size)) - output_shard_sizes[mesh_dim] = tuple(output_shard_sizes[mesh_dim]) - - return_tensor = ShardTensor.from_local( - local_values, - device_mesh=tensor._spec.mesh, - placements=[ - Shard(dim), - ], - sharding_shapes=output_shard_sizes, - ) - return return_tensor - elif index_placement.is_replicate(): - # The output sharding should match the sharding of the original tensor. - output_size = list(tensor.shape) - - # Replace the output size along the indexing dim with the right size: - output_size[dim] = local_values.shape[dim] - # Cast to shard tensor (as replicated, right now): - output = ShardTensor.from_local( - local_values, - device_mesh=tensor._spec.mesh, - placements=[ - Replicate(), - ], - ) - - # Redistribute to the original sharding of the input tensor: - output = output.redistribute(tensor._spec.mesh, tensor._spec.placements) - - return output - else: - raise MissingShardPatch( - f"Index select is not implemented for {index_placement} sharding." - ) - - @staticmethod - def setup_context(ctx, inputs, output) -> None: - r"""Save the source ShardTensorSpec, local shape, dim, and index for backward. - - ``DisableTorchFunctionSubclass`` shielding avoids re-entering the - ShardTensor ``__torch_function__`` fallback while reading - ``tensor._spec`` / ``tensor._local_tensor`` -- the same AOT-hostile - bridge motivated the shielding in ``ShardedSum.setup_context``. - """ - tensor, dim, index = inputs - with torch._C.DisableTorchFunctionSubclass(): - ctx.spec = tensor._spec - ctx.grad_shape = tensor._local_tensor.shape - ctx.dim = dim - ctx.save_for_backward(index) - - @staticmethod - def backward( - ctx: torch.autograd.function.FunctionCtx, grad_output: ShardTensor - ) -> tuple[ShardTensor, None, None]: - r"""Backward pass for the index_select operation on ShardTensors. - - The backward pass sends gradients appropriately to the input tensor. - Therefore, its sharding should match the input tensor's sharding. - - Parameters - ---------- - ctx : torch.autograd.function.FunctionCtx - Context object containing saved tensors and attributes from forward pass. - grad_output : ShardTensor - Gradient of the loss with respect to the output of forward pass. - - Returns - ------- - Tuple[ShardTensor, None, None] - Tuple containing: - - - Gradient with respect to input tensor - - ``None`` for dim parameter (not differentiable) - - ``None`` for index parameter (not differentiable) - """ - (index,) = ctx.saved_tensors - spec = ctx.spec - dim = ctx.dim - - local_index = index.full_tensor() - - grad_inputs = torch.zeros( - spec.tensor_meta.shape, - device=grad_output._local_tensor.device, - dtype=grad_output._local_tensor.dtype, - ) - # local_grad_output = grad_output.to_local() - local_grad_output = grad_output.full_tensor() - - grad_inputs = aten.index_add(grad_inputs, dim, local_index, local_grad_output) - - # Now, grad_inputs is replicated on all devices. - # Shard it along the original sharding of the input tensor. - grad_inputs = ShardTensor.from_local( - grad_inputs, - device_mesh=spec.mesh, - placements=[ - Replicate(), - ], +def _owner_of(index_flat: torch.Tensor, offsets: list[int]) -> torch.Tensor: + r"""Owning rank of every global position. + + Parameters + ---------- + index_flat : torch.Tensor + 1-D non-negative global positions. + offsets : list[int] + The ``W + 1`` shard boundaries along the indexed dim. + + Returns + ------- + torch.Tensor + Rank index in ``[0, W)`` per position. + """ + bounds = torch.tensor(offsets[1:-1], dtype=torch.int64, device=index_flat.device) + return torch.bucketize(index_flat, bounds, right=True) + + +def _rank_request_slice(n_requests: int, rank: int, world_size: int) -> tuple[int, int]: + r"""This rank's contiguous share of ``n_requests`` replicated requests. + + Used in the backward of a gather with a *replicated* index: every rank + holds the full output gradient, so each is responsible for its share of + the requests or the owners would accumulate ``world_size`` copies. + + Parameters + ---------- + n_requests : int + Total number of (flattened) requests. + rank : int + This rank's position on the mesh dim. + world_size : int + Mesh dim size. + + Returns + ------- + tuple[int, int] + ``(lo, hi)`` into the flattened requests; shares differ by at most one. + """ + base, extra = divmod(n_requests, world_size) + lo = rank * base + min(rank, extra) + return lo, lo + base + (1 if rank < extra else 0) + + +@torch.library.custom_op("physicsnemo::exchange_requests", mutates_args=()) +def _exchange_requests_op( + index_flat: torch.Tensor, request_sizes: list[int], group_name: str +) -> torch.Tensor: + r"""Send this rank's requests to every rank; receive every rank's requests. + + Runs once per gather in the forward; the result is saved for the backward, + which therefore needs only the gradient exchange. + + Parameters + ---------- + index_flat : torch.Tensor + This rank's flattened, non-negative global positions. + request_sizes : list[int] + Number of requests held by each rank (spec constants). + group_name : str + c10d group name of the 1-D mesh. + + Returns + ------- + torch.Tensor + All requests, rank-ordered: ``sum(request_sizes)`` positions. + """ + world_size = len(request_sizes) + n_mine = index_flat.numel() + return funcol_all_to_all_v_rows( + index_flat.repeat(world_size), [n_mine] * world_size, request_sizes, group_name + ) + + +@_exchange_requests_op.register_fake +def _exchange_requests_fake( + index_flat: torch.Tensor, request_sizes: list[int], group_name: str +) -> torch.Tensor: + r"""Fake (meta) implementation: ``sum(request_sizes)`` positions.""" + return index_flat.new_empty((sum(request_sizes),)) + + +def _owned_rows( + values: torch.Tensor, + all_requests: torch.Tensor, + offsets: list[int], + rank: int, +) -> tuple[torch.Tensor, torch.Tensor]: + r"""Local row ids of the requests this rank owns, plus the ownership mask. + + Parameters + ---------- + values : torch.Tensor + This rank's shard of rows. + all_requests : torch.Tensor + Every rank's requests (global positions), rank-ordered. + offsets : list[int] + The ``W + 1`` shard boundaries. + rank : int + This rank. + + Returns + ------- + tuple[torch.Tensor, torch.Tensor] + ``(local_rows, owned)``: local row ids clamped into range (filler for + rows this rank does not own) and the boolean mask of owned requests. + """ + lo, hi = offsets[rank], offsets[rank + 1] + owned = (all_requests >= lo) & (all_requests < hi) + n_local = values.shape[0] + local_rows = (all_requests - lo).clamp_(0, max(n_local - 1, 0)) + return local_rows, owned + + +@torch.library.custom_op("physicsnemo::routed_gather", mutates_args=()) +def _routed_gather_op( + values: torch.Tensor, + index: torch.Tensor, + all_requests: torch.Tensor, + offsets: list[int], + request_sizes: list[int], + rank: int, + group_name: str, + index_replicated: bool, +) -> torch.Tensor: + r"""``full_values[index]`` along dim 0 without materializing ``full_values``. + + Parameters + ---------- + values : torch.Tensor + This rank's shard of rows ``[offsets[rank], offsets[rank + 1])``. + index : torch.Tensor + Integer tensor of any shape holding global row ids (already wrapped + to non-negative positions). + all_requests : torch.Tensor + Every rank's requests, rank-ordered (from + :func:`_exchange_requests_op`). + offsets : list[int] + The ``W + 1`` shard boundaries along dim 0. + request_sizes : list[int] + ``index.numel()`` on every rank, in rank order. + rank : int + This rank's position on the mesh dim. + group_name : str + c10d group name of the 1-D mesh. + index_replicated : bool + Whether every rank holds the same requests; unused in the forward, it + rides along so the backward knows to split responsibility. + + Returns + ------- + torch.Tensor + Gathered rows, shape ``(*index.shape, *values.shape[1:])``. + """ + world_size = len(request_sizes) + n_mine = index.numel() + trailing = values.shape[1:] + index_flat = index.reshape(-1) + + # Rows for the requests this rank owns (filler for the rest), then one + # full-length buffer back to every requester. + local_rows, _ = _owned_rows(values, all_requests, offsets, rank) + if values.shape[0] == 0: + rows = values.new_zeros((all_requests.numel(), *trailing)) + else: + rows = values.index_select(0, local_rows) + received = funcol_all_to_all_v_rows( + rows, request_sizes, [n_mine] * world_size, group_name + ) + + # ``received`` holds W buffers of my n_mine requests; take each row from + # its owner. + received = received.reshape(world_size, n_mine, *trailing) + owner = _owner_of(index_flat, offsets) + out = received[owner, torch.arange(n_mine, device=index.device)] + return out.reshape((*index.shape, *trailing)) + + +@_routed_gather_op.register_fake +def _routed_gather_fake( + values: torch.Tensor, + index: torch.Tensor, + all_requests: torch.Tensor, + offsets: list[int], + request_sizes: list[int], + rank: int, + group_name: str, + index_replicated: bool, +) -> torch.Tensor: + r"""Fake (meta) implementation: the output shape is static, no data needed.""" + return values.new_empty((*index.shape, *values.shape[1:])) + + +@torch.library.custom_op("physicsnemo::routed_scatter_add", mutates_args=()) +def _routed_scatter_add_op( + grad: torch.Tensor, + index: torch.Tensor, + all_requests: torch.Tensor, + offsets: list[int], + request_sizes: list[int], + request_slices: list[int], + rank: int, + n_local_rows: int, + group_name: str, +) -> torch.Tensor: + r"""Adjoint of :func:`_routed_gather_op`. + + Every requester sends its full gradient buffer to every owner; each owner + accumulates the rows it owns into a zero tensor of its local rows. + + Parameters + ---------- + grad : torch.Tensor + Gradient of the gathered rows, shape ``(*index.shape, *trailing)``. + index : torch.Tensor + The forward index (non-negative global row ids). + all_requests : torch.Tensor + Every rank's requests, saved from the forward. + offsets : list[int] + The ``W + 1`` shard boundaries along dim 0. + request_sizes : list[int] + ``index.numel()`` on every rank, in rank order. + request_slices : list[int] + ``2 * W`` ints, ``(lo, hi)`` per requester: the share of that + requester's flattened requests whose gradient counts. The full range + for a rank-local (sharded) index; disjoint shares for a replicated + index, where every rank holds the same full gradient. + rank : int + This rank's position on the mesh dim. + n_local_rows : int + Number of rows in this rank's shard of ``values``. + group_name : str + c10d group name of the 1-D mesh. + + Returns + ------- + torch.Tensor + Gradient of this rank's shard of ``values``, shape + ``(n_local_rows, *trailing)``, in ``grad.dtype``. + """ + world_size = len(request_sizes) + n_mine = index.numel() + trailing = grad.shape[index.ndim :] + + grad_rows = grad.reshape(n_mine, *trailing) + incoming = funcol_all_to_all_v_rows( + grad_rows.repeat(world_size, *([1] * len(trailing))), + [n_mine] * world_size, + request_sizes, + group_name, + ) + + # Keep a request iff this rank owns its row and it falls in the sender's + # responsible share. + local_rows, owned = _owned_rows(grad, all_requests, offsets, rank) + position = torch.cat( + [torch.arange(n, device=grad.device, dtype=torch.int64) for n in request_sizes] + ) + lo = torch.tensor( + [request_slices[2 * r] for r in range(world_size)], + device=grad.device, + dtype=torch.int64, + ).repeat_interleave(torch.tensor(request_sizes, device=grad.device)) + hi = torch.tensor( + [request_slices[2 * r + 1] for r in range(world_size)], + device=grad.device, + dtype=torch.int64, + ).repeat_interleave(torch.tensor(request_sizes, device=grad.device)) + keep = owned & (position >= lo) & (position < hi) + + acc_dtype = _accumulator_dtype(grad.dtype) + out = torch.zeros((n_local_rows, *trailing), dtype=acc_dtype, device=grad.device) + weight = keep.to(acc_dtype).reshape(-1, *([1] * len(trailing))) + out.index_add_(0, local_rows, incoming.to(acc_dtype) * weight) + return out.to(grad.dtype) + + +@_routed_scatter_add_op.register_fake +def _routed_scatter_add_fake( + grad: torch.Tensor, + index: torch.Tensor, + all_requests: torch.Tensor, + offsets: list[int], + request_sizes: list[int], + request_slices: list[int], + rank: int, + n_local_rows: int, + group_name: str, +) -> torch.Tensor: + r"""Fake (meta) implementation: the local gradient shape is static.""" + return grad.new_empty((n_local_rows, *grad.shape[index.ndim :])) + + +def _routed_gather_setup_context(ctx, inputs, output) -> None: + r"""Save what the backward needs: index, exchanged requests, spec constants.""" + ( + values, + index, + all_requests, + offsets, + request_sizes, + rank, + group_name, + index_replicated, + ) = inputs + ctx.save_for_backward(index, all_requests) + ctx.index_replicated = index_replicated + ctx.offsets = list(offsets) + ctx.request_sizes = list(request_sizes) + ctx.rank = rank + ctx.n_local_rows = values.shape[0] + ctx.group_name = group_name + + +def _routed_gather_backward(ctx, grad): + r"""Backward of the routed gather: send ``grad`` to the owners, accumulate. + + The request exchange is not repeated: ``all_requests`` was saved from the + forward, so this is a single collective. + """ + index, all_requests = ctx.saved_tensors + world_size = len(ctx.request_sizes) + if ctx.index_replicated: + slices = [ + b + for r in range(world_size) + for b in _rank_request_slice(index.numel(), r, world_size) + ] + else: + slices = [b for n in ctx.request_sizes for b in (0, n)] + grad_values = _routed_scatter_add_op( + grad.contiguous(), + index, + all_requests, + ctx.offsets, + ctx.request_sizes, + slices, + ctx.rank, + ctx.n_local_rows, + ctx.group_name, + ) + return grad_values, None, None, None, None, None, None, None + + +_routed_gather_op.register_autograd( + _routed_gather_backward, setup_context=_routed_gather_setup_context +) + + +def _local_index(index: Any) -> torch.Tensor: + r"""Local rows of an index operand. + + Always via ``to_local()``: dynamo traces that as an op, whereas a raw + ``_local_tensor`` read on a traced subclass resolves to the *real* + tensor and fails fakeification when that tensor is a view. + + Parameters + ---------- + index : ShardTensor or DTensor or torch.Tensor + Index operand. + + Returns + ------- + torch.Tensor + The local tensor of a distributed index, or the plain tensor itself. + """ + if isinstance(index, (ShardTensor, DTensor)): + return index.to_local() + return index + + +def _index_is_sharded(index: Any) -> bool: + r"""Whether *index* is a ``Shard(0)`` ShardTensor (rank-local requests). + + Parameters + ---------- + index : Any + Index operand. + + Returns + ------- + bool + ``True`` for a ``Shard(0)`` ShardTensor, ``False`` for anything + replicated or plain. + + Raises + ------ + MissingShardPatch + If the index is sharded on any other dim, or is Partial. + """ + if not isinstance(index, ShardTensor): + return False + placement = index._spec.placements[0] + if placement.is_partial(): + raise MissingShardPatch("a Partial index is not supported") + if placement.is_shard() and placement.dim != 0: + raise MissingShardPatch( + f"an index sharded on dim {placement.dim} is not supported; " + "shard the index on dim 0 or replicate it" ) - grad_inputs = grad_inputs.redistribute(spec.mesh, spec.placements) + return placement.is_shard() - return grad_inputs, None, None +def _check_index_dtype(index: Any) -> None: + r"""Integer indices only (``int32`` / ``int64``). -def sharded_index_select( - tensor: ShardTensor, - dim: int, - index: ShardTensor, -) -> ShardTensor: - r"""Perform an index_select operation on ShardTensors with autograd support. + Parameters + ---------- + index : Any + Index operand. + + Raises + ------ + MissingShardPatch + If the index dtype is not ``int32`` or ``int64``. + """ + dtype = getattr(index, "dtype", None) + if dtype not in _INDEX_DTYPES: + raise MissingShardPatch(f"index must be int32 or int64, got {dtype}") - This is a thin wrapper around the ShardedIndexSelect autograd function - to make the operation differentiable. + +def _shard_sizes(spec: ShardTensorSpec, tensor_dim: int) -> list[int]: + r"""Per-rank extents of a spec along one tensor dim. + + Parameters + ---------- + spec : ShardTensorSpec + Spec of a tensor on a 1-D mesh. + tensor_dim : int + Tensor dim to read the extents of. + + Returns + ------- + list[int] + One extent per rank, in rank order. + """ + return [s[tensor_dim] for s in spec.sharding_shapes(0)] + + +def routed_gather(values: ShardTensor, index: Any, dim: int) -> ShardTensor: + r"""Differentiable ``values.index_select(dim, index)`` for ``values`` sharded on *dim*. + + Parameters + ---------- + values : ShardTensor + Source tensor, ``Shard(dim)`` on a 1-D device mesh. + index : ShardTensor or DTensor or torch.Tensor + Integer tensor of any shape holding *global* positions along *dim*. + A ``Shard(0)`` ShardTensor is rank-local requests; a replicated + ShardTensor / DTensor / plain tensor is the same requests on every + rank. Negative positions wrap. + dim : int + Dimension of *values* to gather along. + + Returns + ------- + ShardTensor + For a sharded index, ``Shard(dim)`` with the index's per-rank sizes; + for a replicated index, ``Replicate`` (every rank holds the full + result). The selected axis takes the place of *dim* and the remaining + index dims are inserted there, matching + ``full_values.movedim(dim, 0)[index].movedim(...)``; for a 1-D index + this is exactly ``torch.index_select``. + + Raises + ------ + MissingShardPatch + On a multi-dimensional device mesh, a non-integer index, or an index + sharded on a dim other than 0. + """ + spec = values._spec + if spec.mesh.ndim != 1: + raise MissingShardPatch("routed gather supports 1-D device meshes only") + _check_index_dtype(index) + index_sharded = _index_is_sharded(index) + # Plain-Python spec constants that ride into the custom op as ``int[]``: + # shard boundaries of ``values`` and the request count on every rank. No + # tensor is created and nothing is read back from the device. + world_size = spec.mesh.size(0) + offsets = [0, *itertools.accumulate(_shard_sizes(spec, dim))] + idx_local = _local_index(index) + if index_sharded: + request_sizes = [math.prod(shape) for shape in index._spec.sharding_shapes(0)] + else: + request_sizes = [idx_local.numel()] * world_size + + # Global positions, wrapped once; exchanged once and reused by the backward. + idx_local = _wrap_negative(idx_local.to(torch.int64), offsets[-1]) + group_name = resolve_group_name(spec.mesh) + all_requests = _exchange_requests_op( + idx_local.reshape(-1), request_sizes, group_name + ) + + # to_local / from_local are the differentiable ShardTensor <-> local bridges. + local = values.to_local().movedim(dim, 0) + out = _routed_gather_op( + local, + idx_local, + all_requests, + offsets, + request_sizes, + spec.mesh.get_local_rank(), + group_name, + not index_sharded, + ) + # out: (*idx_local.shape, *trailing) with the selected axis first; put it at + # dim. Materialize the permutation so the wrapped local owns its storage + # (a view inside a ShardTensor does not survive dynamo fakeification). + if dim != 0: + out = out.movedim( + tuple(range(idx_local.ndim)), tuple(range(dim, dim + idx_local.ndim)) + ).contiguous() + + global_shape = list(spec.tensor_meta.shape) + if index_sharded: + idx_sizes = _shard_sizes(index._spec, 0) + idx_trailing = list(idx_local.shape[1:]) + shard_shapes = tuple( + tuple(global_shape[:dim] + [n, *idx_trailing] + global_shape[dim + 1 :]) + for n in idx_sizes + ) + return ShardTensor.from_local( + out, spec.mesh, (Shard(dim),), sharding_shapes={0: shard_shapes} + ) + # Explicit (empty) shard shapes: never ``"infer"`` on a traced path. + return ShardTensor.from_local(out, spec.mesh, (Replicate(),), sharding_shapes={}) + + +def sharded_index_select(tensor: ShardTensor, dim: int, index: Any) -> ShardTensor: + r"""``torch.index_select`` on a ShardTensor. Parameters ---------- tensor : ShardTensor - Input tensor to select from. + Source tensor on a 1-D device mesh. dim : int - Dimension along which to index. - index : ShardTensor - Indices to select. + Dimension to select along (negative values wrap). + index : ShardTensor or DTensor or torch.Tensor + 1-D integer index of global positions along *dim*. Returns ------- ShardTensor - Output tensor containing the selected elements. + - ``tensor`` sharded on *dim*: routed gather (communication scales with + the rows indexed); output sharded like the index. + - ``tensor`` sharded on another dim: purely local. A sharded index is + first gathered (eager only) so every rank selects the same rows of + its shard; the output keeps ``tensor``'s placement. + - ``tensor`` replicated: local select; a sharded index yields + ``Shard(dim)`` (each rank selects its own rows), otherwise + ``Replicate``. + + Raises + ------ + MissingShardPatch + On a multi-dimensional mesh, a Partial source, a non-integer index, an + index sharded on a dim other than 0, or a sharded index on an + off-shard select under ``torch.compile`` (needs a collective). """ - return ShardedIndexSelect.apply(tensor, dim, index) + if dim < 0: + dim += tensor.ndim + spec = tensor._spec + if spec.mesh.ndim != 1: + raise MissingShardPatch("index_select supports 1-D device meshes only") + placement = spec.placements[0] + if placement.is_partial(): + raise MissingShardPatch( + "index_select on a Partial ShardTensor is not supported" + ) + _check_index_dtype(index) + index_sharded = _index_is_sharded(index) + + if placement.is_shard() and placement.dim == dim: + return routed_gather(tensor, index, dim) + + if placement.is_shard(): + if index_sharded: + if _is_tracing((tensor, index)): + raise MissingShardPatch( + "index_select off the sharded dim with a sharded index needs an " + "all-gather of the index; not supported under torch.compile" + ) + idx = index.full_tensor() + else: + idx = _local_index(index) + local = tensor.to_local().index_select(dim, idx) + shapes = tuple( + tuple(s[:dim] + (idx.numel(),) + s[dim + 1 :]) + for s in spec.sharding_shapes(0) + ) + return ShardTensor.from_local( + local, spec.mesh, spec.placements, sharding_shapes={0: shapes} + ) + + # Replicated source. + idx = _local_index(index) + local = tensor.to_local().index_select(dim, idx) + if index_sharded: + g = list(spec.tensor_meta.shape) + shapes = tuple( + tuple(g[:dim] + [n] + g[dim + 1 :]) for n in _shard_sizes(index._spec, 0) + ) + return ShardTensor.from_local( + local, spec.mesh, (Shard(dim),), sharding_shapes={0: shapes} + ) + return ShardTensor.from_local(local, spec.mesh, (Replicate(),), sharding_shapes={}) def index_select_wrapper( @@ -246,32 +721,109 @@ def index_select_wrapper( args: tuple[Any, ...], kwargs: dict[str, Any], ) -> ShardTensor: - r"""Wrapper for index_select operation that handles ShardTensors. + r"""``torch.index_select`` / ``Tensor.index_select`` handler. + + Accepts the positional and keyword spellings of ``(input, dim, index)``. Parameters ---------- func : Callable - The original function being wrapped. + The intercepted function. types : tuple[Any, ...] - Types of the input arguments (unused). + Types involved in the dispatch. args : tuple[Any, ...] - Positional arguments containing (tensor, dim, index). + Positional arguments. kwargs : dict[str, Any] - Keyword arguments (unused). + Keyword arguments. Returns ------- ShardTensor - Output tensor containing the selected elements. - """ - - # Extract the tensor and index from the arguments - tensor, dim, index = args + See :func:`sharded_index_select`. + Raises + ------ + MissingShardPatch + On unexpected extra arguments. + """ + kwargs = dict(kwargs or {}) + params = list(args) + tensor = params.pop(0) if params else kwargs.pop("input") + dim = params.pop(0) if params else kwargs.pop("dim") + index = params.pop(0) if params else kwargs.pop("index") + if params or kwargs: + raise MissingShardPatch( + f"unexpected index_select arguments: {params}, {kwargs}" + ) return sharded_index_select(tensor, dim, index) ShardTensor.register_function_handler(torch.index_select, index_select_wrapper) +ShardTensor.register_function_handler(torch.Tensor.index_select, index_select_wrapper) + + +def _is_int_tensor_key(key: Any) -> bool: + r"""Whether *key* is a single integer tensor (advanced indexing). + + Parameters + ---------- + key : Any + ``__getitem__`` key. + + Returns + ------- + bool + ``True`` for an ``int32`` / ``int64`` tensor of any subclass. + """ + return isinstance(key, torch.Tensor) and key.dtype in _INDEX_DTYPES + + +def getitem_wrapper( + func: Callable, + types: tuple[Any, ...], + args: tuple[Any, ...], + kwargs: dict[str, Any], +) -> Any: + r"""``torch.Tensor.__getitem__`` on a ShardTensor. + + A single integer-tensor key on a ``Shard(0)`` source is advanced indexing + along the sharded dim and takes the routed gather. Every other key (ints, + slices, tuples, masks, other placements) follows the default route, which + mirrors the tail of ``ShardTensor.__torch_function__``: straight to + dispatch under tracing, the DTensor fallback otherwise. + + Parameters + ---------- + func : Callable + The intercepted function. + types : tuple[Any, ...] + Types involved in the dispatch. + args : tuple[Any, ...] + ``(tensor, key)``. + kwargs : dict[str, Any] + Keyword arguments (none expected). + + Returns + ------- + Any + The indexed result; a ShardTensor on the routed path. + """ + if len(args) == 2 and not kwargs: + values, key = args + if ( + isinstance(values, ShardTensor) + and values._spec.mesh.ndim == 1 + and values._spec.placements == (Shard(0),) + and _is_int_tensor_key(key) + ): + return routed_gather(values, key, 0) + if _is_tracing(args, kwargs): + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **kwargs) + return _torch_function_fallback_via_dtensor(func, args, kwargs) + + +ShardTensor.register_function_handler(torch.Tensor.__getitem__, getitem_wrapper) def sharded_select_helper(tensor: ShardTensor, dim: int, index: int) -> ShardTensor: diff --git a/test/ci_tests/interrogate_baseline.txt b/test/ci_tests/interrogate_baseline.txt index 8adbf32e63..f5497fb1cc 100644 --- a/test/ci_tests/interrogate_baseline.txt +++ b/test/ci_tests/interrogate_baseline.txt @@ -563,7 +563,6 @@ physicsnemo/distributed/manager.py:DistributedManager.create_groups_from_config physicsnemo/distributed/utils.py:compute_split_shapes physicsnemo/distributed/utils.py:split_tensor_along_dim physicsnemo/domain_parallel/__init__.py:register_custom_ops -physicsnemo/domain_parallel/shard_utils/__init__.py:register_shard_wrappers physicsnemo/mesh/mesh.py:Mesh.n_cells physicsnemo/mesh/mesh.py:Mesh.n_manifold_dims physicsnemo/mesh/mesh.py:Mesh.n_points diff --git a/test/domain_parallel/ops/test_compile_ops.py b/test/domain_parallel/ops/test_compile_ops.py index aa36480289..58d7a24d5b 100644 --- a/test/domain_parallel/ops/test_compile_ops.py +++ b/test/domain_parallel/ops/test_compile_ops.py @@ -121,7 +121,7 @@ def forward(self, tensor: torch.Tensor) -> torch.Tensor: class IndexSelectWrapper(torch.nn.Module): - r"""``torch.index_select(...)`` on a ShardTensor (exercises ``ShardedIndexSelect``).""" + r"""``torch.index_select(...)`` on a ShardTensor (exercises ``sharded_index_select``).""" def __init__(self, dim: int): super().__init__() @@ -330,10 +330,11 @@ def test_compile_shard_redistribute_2d(distributed_mesh_2d): @pytest.mark.multigpu_static @pytest.mark.timeout(180) def test_compile_sharded_index_select_replicated_index_1d(distributed_mesh): - r"""Compile + backward through ``ShardedIndexSelect`` with a replicated index. + r"""Compile + backward through ``sharded_index_select`` off the sharded dim. - A replicated ``index`` keeps the output sharding aligned with the input, - which is the cheaper / less collective-heavy code path inside the op. + The source is ``Shard(2)`` and the selection is along dim 1, so the op is + purely local: the (small) index is gathered and the local shard is + index-selected; the output keeps the input's placement. """ if not torch.cuda.is_available(): pytest.skip("CUDA is not available") @@ -364,6 +365,62 @@ def test_compile_sharded_index_select_replicated_index_1d(distributed_mesh): _run_compile_fwd_bwd(IndexSelectWrapper(dim=dim), [sharded, sharded_index]) +class GetItemWrapper(torch.nn.Module): + r"""``tensor[index]`` with an integer tensor index (exercises the routed gather).""" + + def forward(self, tensor: torch.Tensor, index: torch.Tensor) -> torch.Tensor: + return tensor[index] + + +@pytest.mark.multigpu_static +@pytest.mark.timeout(180) +@pytest.mark.parametrize("index_placement", [Shard(0), Replicate()]) +def test_compile_getitem_routed_gather_1d(distributed_mesh, index_placement): + r"""Compile + backward through the routed gather (``physicsnemo::routed_gather``). + + ``Shard(0)`` source; the index holds global row ids that reference every + rank and is either rank-local (``Shard(0)``) or replicated. Runs the + compiled module twice and compares grads to eager so guard / recompile + problems and stale specs would surface. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + n_rows, n_index = 291, 97 # uneven on 2, 4 and 8 ranks + + torch.manual_seed(7) + original = torch.rand(n_rows, 4, device=dm.device) + index = torch.randint(0, n_rows, (n_index, 3), device=dm.device) + + reference = original.clone().requires_grad_(True) + _scalar_loss(reference[index]).backward() + + sharded = scatter_tensor( + original, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=True, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(index_placement,), + requires_grad=False, + ) + + torch._dynamo.reset() + compiled = torch.compile( + GetItemWrapper(), backend="aot_eager", fullgraph=True, dynamic=False + ) + for _ in range(2): + sharded.grad = None + _scalar_loss(compiled(sharded, sharded_index)).backward() + torch.testing.assert_close(sharded.grad.full_tensor(), reference.grad) + + @pytest.mark.multigpu_static @pytest.mark.timeout(180) def test_compile_unhalo_padding_1d(distributed_mesh): diff --git a/test/domain_parallel/ops/test_exchange.py b/test/domain_parallel/ops/test_exchange.py new file mode 100644 index 0000000000..ae42841111 --- /dev/null +++ b/test/domain_parallel/ops/test_exchange.py @@ -0,0 +1,125 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023 - 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Tests for the row-exchange primitives in ``shard_utils.exchange`` +(``funcol_all_to_all_v_rows`` and ``resolve_group_name``), the transport +shared by the halo scatter correction and the routed gather. +""" + +import pytest +import torch +import torch.distributed as dist + +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel.shard_utils import exchange +from physicsnemo.domain_parallel.shard_utils.exchange import ( + funcol_all_to_all_v_rows, + resolve_group_name, +) + +from .utils import collective_assert, collective_assert_close + + +def _send_counts(rank: int, world_size: int) -> list[int]: + """Rank ``r`` sends ``r + k`` rows to rank ``k``: every pair differs, some are 0.""" + return [rank + k for k in range(world_size)] + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize("dtype", [torch.float32, torch.int64]) +@pytest.mark.parametrize("trailing", [(), (3,), (2, 2)]) +def test_all_to_all_v_rows_routes_by_rank(distributed_mesh, dtype, trailing): + """Rows land on the destination rank, in source-rank order, with the right shape.""" + dm = DistributedManager() + mesh = distributed_mesh + rank = mesh.get_local_rank(0) + world_size = dist.get_world_size(group=mesh.get_group(0)) + + send_counts = _send_counts(rank, world_size) + # Row values encode (source, destination) so the receiver can verify both. + blocks = [] + for dst, n in enumerate(send_counts): + block = torch.full( + (n, *trailing), rank * 100 + dst, dtype=dtype, device=dm.device + ) + blocks.append(block) + send_rows = torch.cat(blocks, dim=0) + # Rank r receives ``src + r`` rows from every source rank ``src``. + recv_counts = [src + rank for src in range(world_size)] + + received = funcol_all_to_all_v_rows(send_rows, send_counts, recv_counts, mesh) + + collective_assert( + tuple(received.shape) == (sum(recv_counts), *trailing), + msg=f"received shape {tuple(received.shape)}", + ) + expected = torch.cat( + [ + torch.full((n, *trailing), src * 100 + rank, dtype=dtype, device=dm.device) + for src, n in enumerate(recv_counts) + ], + dim=0, + ) + collective_assert_close(received, expected, atol=0, rtol=0, msg="a2a-v rows") + + +@pytest.mark.multigpu_static +def test_all_to_all_v_rows_empty_on_one_rank(distributed_mesh): + """A rank that sends and receives nothing still completes the exchange.""" + dm = DistributedManager() + mesh = distributed_mesh + rank = mesh.get_local_rank(0) + world_size = dist.get_world_size(group=mesh.get_group(0)) + + # Only rank 0 sends: 2 rows to every other rank, nothing to itself. + send_counts = [0] + [2] * (world_size - 1) if rank == 0 else [0] * world_size + recv_counts = [0] * world_size if rank == 0 else [2] + [0] * (world_size - 1) + send_rows = torch.full((sum(send_counts), 4), float(rank), device=dm.device) + + received = funcol_all_to_all_v_rows(send_rows, send_counts, recv_counts, mesh) + + expected = torch.zeros((sum(recv_counts), 4), device=dm.device) + collective_assert(tuple(received.shape) == tuple(expected.shape), msg="shape") + collective_assert_close(received, expected, atol=0, rtol=0, msg="from rank 0") + + +@pytest.mark.multigpu_static +def test_resolve_group_name(distributed_mesh): + """Mesh, process group, string and ``None`` all resolve to a c10d group-name token.""" + mesh = distributed_mesh + group = mesh.get_group(0) + + name = resolve_group_name(mesh) + collective_assert(isinstance(name, str) and name != "", msg="mesh -> name") + collective_assert(resolve_group_name(group) == name, msg="group matches mesh") + collective_assert(resolve_group_name(name) == name, msg="string passthrough") + collective_assert(resolve_group_name(None) == "", msg="None -> default group") + + +def test_fp32_scatter_accumulator_setting(): + """``float32`` scatter folds accumulate in the configured dtype (default float64).""" + assert exchange.get_fp32_scatter_accumulator() is torch.float64 + assert exchange._accumulator_dtype(torch.float32) is torch.float64 + assert exchange._accumulator_dtype(torch.bfloat16) is torch.float32 + assert exchange._accumulator_dtype(torch.int64) is torch.int64 + try: + exchange.set_fp32_scatter_accumulator(torch.float32) + assert exchange._accumulator_dtype(torch.float32) is torch.float32 + with pytest.raises(ValueError): + exchange.set_fp32_scatter_accumulator(torch.float16) + finally: + exchange.set_fp32_scatter_accumulator(torch.float64) diff --git a/test/domain_parallel/ops/test_index.py b/test/domain_parallel/ops/test_index.py new file mode 100644 index 0000000000..a735dea785 --- /dev/null +++ b/test/domain_parallel/ops/test_index.py @@ -0,0 +1,350 @@ +# SPDX-FileCopyrightText: Copyright (c) 2023 - 2026 NVIDIA CORPORATION & AFFILIATES. +# SPDX-FileCopyrightText: All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Test integer-tensor indexing (``tensor[index]``) on ShardTensor. + +The source tensor is ``Shard(0)``; the index holds *global* row ids that +reference every rank, so the gather has to route rows between ranks. The +index is either sharded (``Shard(0)``, each rank asks for its own rows) or +replicated (every rank asks for the same rows). ``torch.index_select`` is +covered in ``test_select.py``. +""" + +import pytest +import torch +from torch.distributed.tensor.placement_types import Replicate, Shard + +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel import ShardTensor, scatter_tensor +from physicsnemo.domain_parallel.shard_utils.patch_core import MissingShardPatch + +from .utils import numerical_shard_tensor_check + +# Row counts that split unevenly on 2, 4 and 8 ranks. +N_ROWS = 291 +N_INDEX = 97 + + +class GetItemWrapper(torch.nn.Module): + """ + Wrapper class for testing ``tensor[index]`` with an integer tensor index. + """ + + def forward(self, tensor: torch.Tensor, index: torch.Tensor): + return tensor[index] + + +def _check_sharded_like_index(index_placement): + def _check(output): + assert isinstance(output, ShardTensor) + assert output._spec.placements == index_placement + + return _check + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize("backward", [False, True]) +def test_getitem_sharded_index(distributed_mesh, backward): + """Sharded source, sharded index: output is sharded like the index.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device, requires_grad=backward) + index = torch.randint(0, N_ROWS, (N_INDEX, 3), device=dm.device) + + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=backward, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + numerical_shard_tensor_check( + distributed_mesh, + GetItemWrapper(), + [sharded_tensor, sharded_index], + {}, + check_grads=backward, + output_check_fn=_check_sharded_like_index((Shard(0),)), + ) + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize("backward", [False, True]) +def test_getitem_replicated_index(distributed_mesh, backward): + """Sharded source, replicated index: every rank holds the full result. + + In backward every rank holds the full output gradient; each routes only + its share of the requests so owners accumulate each row exactly once. + """ + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device, requires_grad=backward) + index = torch.randint(0, N_ROWS, (N_INDEX, 3), device=dm.device) + + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=backward, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Replicate(),), + requires_grad=False, + ) + + numerical_shard_tensor_check( + distributed_mesh, + GetItemWrapper(), + [sharded_tensor, sharded_index], + {}, + check_grads=backward, + output_check_fn=_check_sharded_like_index((Replicate(),)), + ) + + +@pytest.mark.multigpu_static +def test_getitem_bf16(distributed_mesh): + """Reduced-precision source: forward and backward match the local result.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand( + N_ROWS, 4, device=dm.device, dtype=torch.bfloat16, requires_grad=True + ) + index = torch.randint(0, N_ROWS, (N_INDEX, 3), device=dm.device) + + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=True, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + numerical_shard_tensor_check( + distributed_mesh, + GetItemWrapper(), + [sharded_tensor, sharded_index], + {}, + check_grads=True, + atol=1e-2, + rtol=1e-2, + ) + + +@pytest.mark.multigpu_static +def test_getitem_never_all_gathers(distributed_mesh): + """The routed gather issues no all-gather of the source in forward or + backward -- the collective the replaced implementation was built on. + + ``CommDebugMode`` is a dispatch mode and cannot see the all-to-all calls + inside the custom-op bodies, so this pins the absence of the expensive + collective rather than the presence of the cheap one. The only + collective visible at this level is the ``all_reduce`` that resolves the + ``Partial`` result of ``mean`` over the sharded axis. + """ + from torch.distributed.tensor.debug import CommDebugMode + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device) + index = torch.randint(0, N_ROWS, (N_INDEX, 3), device=dm.device) + + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=True, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + with CommDebugMode() as comm: + gathered = sharded_tensor[sharded_index] + gathered.mean().backward() + + counts = {str(op): n for op, n in comm.get_comm_counts().items() if n} + gathers = {op: n for op, n in counts.items() if "all_gather" in op} + # Identical on every rank: the count dict is built from the same program. + assert not gathers, f"routed gather all-gathered: {counts}" + + +@pytest.mark.multigpu_static +def test_getitem_duplicate_and_negative_index(distributed_mesh): + """Repeated and negative global ids: duplicates accumulate in backward, + negatives wrap like eager indexing.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device, requires_grad=True) + # Every rank asks for the same handful of rows several times, some negative. + index = torch.tensor( + [[0, -1, 5], [5, 5, -N_ROWS], [N_ROWS - 1, 0, 0]] * 11, device=dm.device + ) + + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=True, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + numerical_shard_tensor_check( + distributed_mesh, + GetItemWrapper(), + [sharded_tensor, sharded_index], + {}, + check_grads=True, + ) + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize( + "key", + ["int", "slice", "tuple", "mask", "ellipsis"], +) +def test_getitem_non_tensor_keys_fall_through(distributed_mesh, key): + """Keys other than an integer tensor take the default route and match eager.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device) + sharded_tensor = scatter_tensor( + values, global_src=0, mesh=distributed_mesh, placements=(Shard(0),) + ) + mask = torch.zeros(4, dtype=torch.bool, device=dm.device) + mask[1] = mask[3] = True + keys = { + "int": (lambda t: t[3]), + "slice": (lambda t: t[10:50]), + "tuple": (lambda t: t[:, 1:3]), + "mask": (lambda t: t[:, mask]), + "ellipsis": (lambda t: t[..., 0]), + } + out = keys[key](sharded_tensor) + expected = keys[key](values) + out = out.full_tensor() if isinstance(out, ShardTensor) else out + torch.testing.assert_close(out, expected) + + +@pytest.mark.multigpu_static +def test_getitem_rejects_unsupported_index(distributed_mesh): + """Float index, Partial index and an index sharded on dim 1 raise + ``MissingShardPatch`` (communication-free, identical on every rank).""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device) + sharded_tensor = scatter_tensor( + values, global_src=0, mesh=distributed_mesh, placements=(Shard(0),) + ) + world = distributed_mesh.size(0) + + float_index = ShardTensor.from_local( + torch.zeros(3, 2, device=dm.device), + distributed_mesh, + (Shard(0),), + sharding_shapes={0: [(3, 2)] * world}, + ) + with pytest.raises(MissingShardPatch): + sharded_tensor[float_index] + + dim1_index = ShardTensor.from_local( + torch.zeros(3, 2, dtype=torch.int64, device=dm.device), + distributed_mesh, + (Shard(1),), + sharding_shapes={0: [(3, 2)] * world}, + ) + with pytest.raises(MissingShardPatch): + sharded_tensor[dim1_index] + + +@pytest.mark.multigpu_static +def test_getitem_rejects_2d_mesh(distributed_mesh_2d): + """The routed gather is 1-D mesh only.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + torch.manual_seed(7) + values = torch.rand(N_ROWS, 4, device=dm.device) + sharded_tensor = scatter_tensor( + values, + global_src=0, + mesh=distributed_mesh_2d, + placements=(Shard(0), Replicate()), + ) + index = torch.randint(0, N_ROWS, (5,), device=dm.device) + with pytest.raises(MissingShardPatch): + torch.index_select(sharded_tensor, 0, index) diff --git a/test/domain_parallel/ops/test_select.py b/test/domain_parallel/ops/test_select.py index 891dbde96a..f675221f96 100644 --- a/test/domain_parallel/ops/test_select.py +++ b/test/domain_parallel/ops/test_select.py @@ -17,15 +17,15 @@ """ Test selection operations on ShardTensor. This file tests both torch.select and torch.index_select. We use a 3D tensor to -do the tests, it has no special significance. We're not testing -over all possible selection dimensions, especially not along sharded dimensions. +do the tests, it has no special significance. -That could be implemented in the future. +``index_select`` is covered both off the sharded dimension (a purely local +op) and along it (the routed gather, with a sharded or a replicated index). """ import pytest import torch -from torch.distributed.tensor.placement_types import Shard +from torch.distributed.tensor.placement_types import Replicate, Shard from physicsnemo.distributed import DistributedManager from physicsnemo.domain_parallel import scatter_tensor @@ -66,7 +66,7 @@ def test_select_operation( distributed_mesh, backward, ): - """Test basic scaled dot product attention with various configurations""" + """``torch.select`` on a ``Shard(2)`` tensor along an unsharded dim (local op).""" if not torch.cuda.is_available(): pytest.skip("CUDA is not available") @@ -108,7 +108,8 @@ def test_index_select_operation( distributed_mesh, backward, ): - """Test basic scaled dot product attention with various configurations""" + """``index_select`` off the sharded dim: the index is gathered, the shard is + selected locally, and the output keeps the input placement.""" if not torch.cuda.is_available(): pytest.skip("CUDA is not available") @@ -150,3 +151,154 @@ def test_index_select_operation( {}, check_grads=backward, ) + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize("target_dim", [0, 1]) +@pytest.mark.parametrize("backward", [False, True]) +def test_index_select_along_sharded_dim( + distributed_mesh, + target_dim, + backward, +): + """``index_select`` along the sharded dim with a sharded index (routed gather). + + The index holds global positions along the sharded dim that reference + every rank; the output is sharded like the index. + """ + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + shape = (61, 131, 8) # uneven on 2, 4 and 8 ranks + N = 97 + + original_tensor = torch.rand(shape, device=dm.device, requires_grad=backward) + index = torch.randint(low=0, high=shape[target_dim], size=(N,), device=dm.device) + + sharded_tensor = scatter_tensor( + original_tensor, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(target_dim),), + requires_grad=backward, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + def check_output(output): + assert output._spec.placements == (Shard(target_dim),) + + numerical_shard_tensor_check( + distributed_mesh, + IndexSelectWrapper(target_dim=target_dim), + [sharded_tensor, sharded_index], + {}, + check_grads=backward, + output_check_fn=check_output, + ) + + +@pytest.mark.multigpu_static +@pytest.mark.parametrize("target_dim", [0, 1]) +@pytest.mark.parametrize("backward", [False, True]) +def test_index_select_along_sharded_dim_replicated_index( + distributed_mesh, + target_dim, + backward, +): + """``index_select`` along the sharded dim with a replicated index. + + Every rank asks for the same rows, so every rank holds the full result; + in backward each rank routes only its share of the requests. + """ + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + shape = (61, 131, 8) + N = 97 + + original_tensor = torch.rand(shape, device=dm.device, requires_grad=backward) + index = torch.randint(low=0, high=shape[target_dim], size=(N,), device=dm.device) + + sharded_tensor = scatter_tensor( + original_tensor, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(target_dim),), + requires_grad=backward, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Replicate(),), + requires_grad=False, + ) + + def check_output(output): + assert output._spec.placements == (Replicate(),) + + numerical_shard_tensor_check( + distributed_mesh, + IndexSelectWrapper(target_dim=target_dim), + [sharded_tensor, sharded_index], + {}, + check_grads=backward, + output_check_fn=check_output, + ) + + +class IndexSelectMethodWrapper(torch.nn.Module): + """``tensor.index_select(dim=..., index=...)``: method spelling with keywords.""" + + def __init__(self, target_dim: int): + super().__init__() + self.target_dim = target_dim + + def forward(self, tensor: torch.Tensor, index: torch.Tensor): + return tensor.index_select(dim=self.target_dim, index=index) + + +@pytest.mark.multigpu_static +def test_index_select_method_spelling(distributed_mesh): + """``Tensor.index_select`` with keyword arguments takes the same handler.""" + + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + shape = (61, 8) + original_tensor = torch.rand(shape, device=dm.device, requires_grad=True) + index = torch.randint(low=0, high=shape[0], size=(97,), device=dm.device) + + sharded_tensor = scatter_tensor( + original_tensor, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=True, + ) + sharded_index = scatter_tensor( + index, + global_src=0, + mesh=distributed_mesh, + placements=(Shard(0),), + requires_grad=False, + ) + + numerical_shard_tensor_check( + distributed_mesh, + IndexSelectMethodWrapper(target_dim=0), + [sharded_tensor, sharded_index], + {}, + check_grads=True, + )