From 9831c246e117f093e111902b9f116c7d445dbbe8 Mon Sep 17 00:00:00 2001 From: Corey Adams <6619961+coreyjadams@users.noreply.github.com> Date: Tue, 8 Sep 2026 21:26:48 +0000 Subject: [PATCH] Routed gather for sharded index_select and integer-tensor indexing Replace the all-gathering ShardedIndexSelect with a routed exchange: a Shard(d) source and an index of global positions (sharded or replicated) move only the requested rows between ranks. All exchange sizes come from the specs, so the custom ops have no device-to-host sync and trace under torch.compile; the backward reuses the forward's request exchange and is a single collective. - shard_utils/exchange.py (new): row all-to-all, group-name resolution and the fp32 scatter accumulator policy (float64 default, configurable), shared by halo_scatter and index_ops. - index_ops.py: routed_gather / sharded_index_select; handlers for torch.index_select, Tensor.index_select and Tensor.__getitem__ with a single int-tensor key; replicated-index backward splits responsibility across ranks; negative indices wrap; 1-D mesh and int dtype guards. - shard_tensor.py: _ToTorchTensor.backward adopts the gradient's shard order with its placements. - Tests: test_index.py (getitem), index_select along the sharded dim in test_select.py, test_exchange.py, routed-gather compile test. --- CHANGELOG.md | 6 + physicsnemo/domain_parallel/shard_tensor.py | 6 +- .../domain_parallel/shard_utils/__init__.py | 1 + .../domain_parallel/shard_utils/exchange.py | 193 ++++ .../shard_utils/halo_scatter.py | 70 +- .../domain_parallel/shard_utils/index_ops.py | 936 ++++++++++++++---- test/ci_tests/interrogate_baseline.txt | 1 - test/domain_parallel/ops/test_compile_ops.py | 65 +- test/domain_parallel/ops/test_exchange.py | 125 +++ test/domain_parallel/ops/test_index.py | 350 +++++++ test/domain_parallel/ops/test_select.py | 164 ++- 11 files changed, 1654 insertions(+), 263 deletions(-) create mode 100644 physicsnemo/domain_parallel/shard_utils/exchange.py create mode 100644 test/domain_parallel/ops/test_exchange.py create mode 100644 test/domain_parallel/ops/test_index.py 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, + )