Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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

Expand Down
6 changes: 5 additions & 1 deletion physicsnemo/domain_parallel/shard_tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

shard_order is an upstream attribute we're simply adding in here.

tensor_meta=grad_output._spec.tensor_meta,
)
return (_dtensor_to_shard_tensor(grad_output, cached_spec),)
Expand Down
1 change: 1 addition & 0 deletions physicsnemo/domain_parallel/shard_utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
193 changes: 193 additions & 0 deletions physicsnemo/domain_parallel/shard_utils/exchange.py
Original file line number Diff line number Diff line change
@@ -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",
]
Comment on lines +34 to +39

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This file, exchange.py, is a brand new file that consolidates the halo exchange work supporting alchemi merged in v2.2. It was useful and necessary to reproduce almost all of that for index and index_select handlers properly, so I moved core functionality here to mirror halo.py and ring.py as a slightly-lower-level utility than ShardTensor. We could do optimizations on these functions and it would now hit all paths, etc.



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)

Check warning on line 58 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L58

Added line #L58 was not covered by tests
if group is None:
return dist.distributed_c10d._get_default_group()
return group

Check warning on line 61 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L60-L61

Added lines #L60 - L61 were not covered by tests


# 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 ""

Check warning on line 146 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L146

Added line #L146 was not covered by tests
if isinstance(group, str):
return group

Check warning on line 148 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L148

Added line #L148 was not covered by tests
if isinstance(group, DeviceMesh):
return group.get_group(0).group_name
return group.group_name

Check warning on line 151 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L150-L151

Added lines #L150 - L151 were not covered by tests


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

Check warning on line 181 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L180-L181

Added lines #L180 - L181 were not covered by tests
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(

Check warning on line 188 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L183-L188

Added lines #L183 - L188 were not covered by tests
funcol.all_to_all_single(
flat_send, recv_flat, send_flat, _funcol_group_arg(group)
)
)
return flat_recv.reshape((total_recv,) + trailing)

Check warning on line 193 in physicsnemo/domain_parallel/shard_utils/exchange.py

View check run for this annotation

Codecov / codecov/patch

physicsnemo/domain_parallel/shard_utils/exchange.py#L193

Added line #L193 was not covered by tests
70 changes: 11 additions & 59 deletions physicsnemo/domain_parallel/shard_utils/halo_scatter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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:
Comment on lines -49 to -53

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The changes to this file should be exclusively moving things to exchange.py without functionality change.

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.

Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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
Expand Down
Loading
Loading