-
Notifications
You must be signed in to change notification settings - Fork 800
Routed gather for sharded index_select and integer-tensor indexing #1981
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This file, |
||
|
|
||
|
|
||
| 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) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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: | ||
|
Comment on lines
-49
to
-53
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The changes to this file should be exclusively moving things to |
||
| 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 | ||
|
|
||
There was a problem hiding this comment.
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.