From 12ae16bd8e7fc6351994ad9e4367619244066cb8 Mon Sep 17 00:00:00 2001 From: Corey Adams <6619961+coreyjadams@users.noreply.github.com> Date: Wed, 9 Sep 2026 22:12:03 +0000 Subject: [PATCH] Domain-parallel (rank-local) reading in the datapipes Readers can produce sharded samples directly: each rank reads only its share of every sharded batch axis from disk, the sample travels as a ShardedProto, and the dataset assembles Shard(0) ShardTensors on the device. Nothing the size of the full sample exists on any rank. - _domain_parallel.py: DomainParallelConfig parsed once from the `domain_parallel` dict + 1-D device_mesh (auto_shard_size gate on the length of dim 0, `placements` overrides by dotted prefix, unmatched keys warn); ShardedProto payload and the communication-free ShardTensor wrap. - Reader: `_load_sample_domain_parallel` contract for custom readers, `_selection_plan` shared by the plain and rank-local reads; ZarrReader and TensorStoreZarrReader implement it with one read procedure each. DomainParallelConfig and resolve_leaf_placements are exported from physicsnemo.datapipes.readers for custom readers. - MeshReader / DomainMeshReader: one batch axis per mesh (cells, with vertices following; points for a point cloud), even splits of the cell and point axes, cells keep global vertex ids so points[cells] is the routed gather; cell windows compact globally as the eager reader does. Memmap and zarr share `_read_mesh_selection`; extra boundaries pin to replicate; subsampling requires a seed. Shared reader base for seed, epoch, zarr group cache and close(). - CenterMesh materializes the center of mass on sharded meshes. - Tests: CPU config/plan tests; GPU datapipe tests for zarr, memmap and zarr meshes, domain meshes (nested global_data, extra boundaries, drop flags, placement overrides, subsample equivalence). Minimal example. User-guide docs follow in a separate PR. --- CHANGELOG.md | 8 + .../datapipes/domain_parallel_reading.py | 158 +++ physicsnemo/datapipes/_domain_parallel.py | 570 +++++++++++ physicsnemo/datapipes/dataset.py | 3 + physicsnemo/datapipes/mesh_dataset.py | 7 + physicsnemo/datapipes/protocols.py | 34 +- physicsnemo/datapipes/readers/__init__.py | 10 + physicsnemo/datapipes/readers/base.py | 199 +++- physicsnemo/datapipes/readers/mesh.py | 907 +++++++++++++++--- .../datapipes/readers/tensorstore_zarr.py | 72 +- physicsnemo/datapipes/readers/zarr.py | 88 +- .../datapipes/transforms/mesh/transforms.py | 12 +- test/datapipes/test_domain_parallel_config.py | 374 ++++++++ test/domain_parallel/datapipes/__init__.py | 15 + .../test_sharded_domain_mesh_dataset.py | 392 ++++++++ .../datapipes/test_sharded_mesh_dataset.py | 284 ++++++ .../test_sharded_tensordict_dataset.py | 204 ++++ 17 files changed, 3182 insertions(+), 155 deletions(-) create mode 100644 examples/minimal/datapipes/domain_parallel_reading.py create mode 100644 physicsnemo/datapipes/_domain_parallel.py create mode 100644 test/datapipes/test_domain_parallel_config.py create mode 100644 test/domain_parallel/datapipes/__init__.py create mode 100644 test/domain_parallel/datapipes/test_sharded_domain_mesh_dataset.py create mode 100644 test/domain_parallel/datapipes/test_sharded_mesh_dataset.py create mode 100644 test/domain_parallel/datapipes/test_sharded_tensordict_dataset.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 16bb833938..f16c710f75 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -26,6 +26,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 collectives; ring attention works under `torch.compile` as an eager graph-break region. In-place `detach_` on a `ShardTensor` is supported. - `radius_search` on sharded inputs works under `torch.compile`. +- Domain-parallel (rank-local) reading in the datapipes: `ZarrReader`, + `TensorStoreZarrReader`, `MeshReader` and `DomainMeshReader` accept a + `domain_parallel` dict and a 1-D `device_mesh`; each rank reads only its + rows and the dataset assembles `Shard(0)` ShardTensors on the device. + Custom readers opt in by implementing `Reader._load_sample_domain_parallel` + with the exported `DomainParallelConfig` / `resolve_leaf_placements`. ### Changed @@ -35,6 +41,8 @@ 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. +- `CenterMesh` materializes the center of mass on sharded meshes so the offset + stays a plain tensor. - `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. diff --git a/examples/minimal/datapipes/domain_parallel_reading.py b/examples/minimal/datapipes/domain_parallel_reading.py new file mode 100644 index 0000000000..04749215ed --- /dev/null +++ b/examples/minimal/datapipes/domain_parallel_reading.py @@ -0,0 +1,158 @@ +# 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"""Minimal domain-parallel (rank-local) reading example. + +Each rank reads only its chunk of the large arrays straight from disk; the +sample arrives as ``Shard(0)`` ShardTensors over a device mesh you construct +in Python and inject into the reader. Small arrays replicate automatically. + +Run with:: + + torchrun --nproc-per-node 2 examples/minimal/datapipes/domain_parallel_reading.py + +The pattern to take away for recipes: + +1. The ``domain_parallel`` dict is plain, Hydra-friendly configuration. +2. The ``DeviceMesh`` is NOT configuration -- construct it at runtime from + ``DistributedManager`` and pass it to the reader alongside the dict. +3. The dataset needs no domain-parallel arguments at all; readers own the + rank-local read, datasets assemble the ShardTensors on the GPU. +4. Every rank of the domain mesh must ask for the same sample index (here + all ranks read ``dataset[0]``); with a data-parallel axis, shard the + sampler over that axis only. Subsampling needs a seed so every rank + draws the same window. +""" + +import shutil +import tempfile +from pathlib import Path + +import numpy as np +import torch +import torch.distributed as dist +import zarr + +from physicsnemo.datapipes import Dataset +from physicsnemo.datapipes.readers.zarr import ZarrReader +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel import ShardTensor + + +def generate_sample_data(root: Path, n_samples: int = 4, n_points: int = 100_000): + r"""Write example zarr groups: two large point-wise arrays, one small one. + + The size split is deliberate: ``coords`` and ``fields`` share a batch + axis long enough to pass the example's ``auto_shard_size`` gate and + will be sharded across the domain mesh, while ``params`` is short and + will replicate. + + Parameters + ---------- + root : Path + Directory to write the ``sample_.zarr`` groups into. + n_samples : int, default=4 + Number of zarr groups (samples) to create. + n_points : int, default=100_000 + Number of rows in the large point-wise arrays. + """ + rng = np.random.default_rng(0) + for i in range(n_samples): + group = zarr.open_group(str(root / f"sample_{i}.zarr"), mode="w") + group["coords"] = rng.standard_normal((n_points, 3), dtype=np.float32) + group["fields"] = rng.standard_normal((n_points, 4), dtype=np.float32) + group["params"] = rng.standard_normal((8,), dtype=np.float32) + + +def main(): + r"""Run the domain-parallel reading example end to end. + + Initializes distributed, builds a 1-D device mesh, generates example + data on rank 0, constructs a ``ZarrReader`` with a declarative + ``domain_parallel`` policy plus the runtime-injected mesh, and prints + each key's global shape, rank-local shape, and placement to show which + arrays were sharded versus replicated. + """ + DistributedManager.initialize() + dm = DistributedManager() + + # The device mesh is a runtime object: build it here, in Python, and + # inject it into the reader. It never appears in yaml/Hydra config. + device_mesh = dm.initialize_mesh([-1], ["domain"]) + + # Rank 0 generates example data in a shared location. + if dm.rank == 0: + root = Path(tempfile.mkdtemp(prefix="dp_datapipe_example_")) + generate_sample_data(root) + holder = [str(root)] + else: + holder = [None] + dist.broadcast_object_list(holder, src=0) + data_root = holder[0] + + reader = ZarrReader( + data_root, + # Optional: coordinated subsampling composes with domain-parallel + # reading -- each rank reads its chunk OF the subsampled window. + coordinated_subsampling={ + "n_points": 50_000, + "target_keys": ["coords", "fields"], + }, + # Declarative policy (Hydra-friendly): a batch axis shards when its + # length (tensor dim 0) is at least this many entries, decided from + # store metadata before any data is read. ``placements`` could pin + # axes explicitly. + domain_parallel={"auto_shard_size": 1024}, + # Runtime object, injected in Python. + device_mesh=device_mesh, + ) + + dataset = Dataset(reader, device=dm.device) + # Subsampling draws a window per sample; the seed makes it identical on + # every rank (the DataLoader does this for you when given a seed). + generator = torch.Generator() + generator.manual_seed(0) + dataset.set_generator(generator) + + sample, metadata = dataset[0] + lines = [] + for key, value in sample.items(): + if isinstance(value, ShardTensor): + local = value.to_local() + placement = value.placements + else: + local, placement = value, "(replicated)" + lines.append( + f" rank {dm.rank}: {key}: global {tuple(value.shape)}, " + f"local {tuple(local.shape)}, {placement}" + ) + # Print rank by rank so the report reads cleanly. + for rank in range(dm.world_size): + if rank == dm.rank: + if rank == 0: + print(f"sample from {metadata['source_filename']}:") + print("\n".join(lines), flush=True) + dist.barrier() + + dataset.close() + dist.barrier() + if dm.rank == 0: + shutil.rmtree(data_root, ignore_errors=True) + DistributedManager.cleanup() + + +if __name__ == "__main__": + main() diff --git a/physicsnemo/datapipes/_domain_parallel.py b/physicsnemo/datapipes/_domain_parallel.py new file mode 100644 index 0000000000..f52d42676f --- /dev/null +++ b/physicsnemo/datapipes/_domain_parallel.py @@ -0,0 +1,570 @@ +# 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"""Shared machinery for domain-parallel (rank-local) reading in datapipes. + +Readers that support domain-parallel reading accept a ``domain_parallel`` +configuration dict plus a 1-D ``device_mesh`` (constructed and injected at +runtime in Python -- it is not serializable config). Each rank reads only +its share of the sharded batch axes; the local pieces move to the GPU inside a +:class:`ShardedProto`; and :meth:`ShardedProto.assemble` builds the sample +with ``Shard(0)`` ShardTensors via the communication-free chunk path of +``ShardTensor.from_local``. + +Configuration +------------- +:: + + domain_parallel = { + # Auto gate: shard a batch axis when its length (the size of tensor + # dim 0) is at least this many entries; shorter axes replicate. + "auto_shard_size": 1024, + # Optional overrides, keyed by axis name: "shard" | "replicate". + # Anything not named falls back to the auto gate. + "placements": {"interior.points": "shard", "boundaries.stl": "replicate"}, + } + +``auto_shard_size`` looks at the length of dim 0 for each tensor, not the +total number of elements. With the default of 1024: a point cloud of shape +``(200_000, 3)`` shards, while a ``(512, 512)`` image and a ``(7, 100000)`` +table replicate (pin them with ``placements`` to shard them along dim 0). An +axis shorter than the world size always replicates, so no rank is left with +an empty shard. ``auto_shard_size: 1`` shards everything except scalars and +axes shorter than the world size. + +Decisions are made per **batch axis**, not per tensor. Tensors that share a +batch axis are co-indexed and placed together. A mesh has one batch axis: +``cells`` when it has cells (``cells``, every ``cell_data`` leaf, and its +``points`` / ``point_data`` all chunk together; ``cells`` keeps global vertex +ids and ``points[cells]`` is the routed gather), or ``points`` for a point +cloud. For a flat dataset the axes are the groups of leaves that share a +dim-0 length. For example, a zarr dataset with ``coords`` and ``fields`` +arrays of the same length shards both even if you only specify +``placements: {"coords": "shard"}``. Subsampling often forces the same dim-0 +length across a dataset, and the gate is decided *after* subsampling: the +axis length it sees is the subsampled length, not the stored one. + +Axis names are dotted paths. An override applies to the axis it names and, +by prefix, to every axis beneath it: ``boundaries.stl: replicate`` pins both +``boundaries.stl.points`` and ``boundaries.stl.cells``; for a flat sample +``solution: shard`` pins every leaf under ``solution``. Overrides that match +nothing are reported as warnings. + +Placement convention +-------------------- +A sharded leaf is a ``ShardTensor`` with ``Shard(0)``. A replicated leaf is a +**plain tensor** that is identical on every rank of the device mesh -- it is +never wrapped as a ``Replicate`` ShardTensor. Mixed ops rely on ShardTensor +promoting plain operands, and reductions over a sharded axis yield +``Partial`` results that callers materialize (``full_tensor``) only where a +rank-identical value is required. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass, field, replace +from typing import Any, Callable, Iterable, Literal, Mapping + +import torch +from tensordict import TensorDict +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.tensor.placement_types import Shard + +from physicsnemo.datapipes.keys import NestedKey, key_to_str +from physicsnemo.domain_parallel import ShardTensor +from physicsnemo.domain_parallel._shard_tensor_spec import ( + compute_sharding_shapes_from_chunking_global_shape, +) + +logger = logging.getLogger(__name__) + +# The only placement domain-parallel readers produce: sharded on the batch +# (dim-0) axis of a 1-D device mesh. +PLACEMENTS = (Shard(0),) + +# A sharded tensor pays a fixed per-op dispatch cost regardless of size, so +# short batch axes are strictly cheaper replicated. This default keeps +# per-sample metadata (freestream vectors, parameter tables) plain while +# sharding anything point- or cell-sized. +DEFAULT_AUTO_SHARD_SIZE = 1024 + +# Placement names as they appear in configuration. The ``domain_parallel`` +# dict is parsed from YAML (Hydra), so overrides are strings rather than +# ``torch.distributed`` ``Placement`` objects; they are promoted to the real +# placement at assembly time, where the only outcomes are ``Shard(0)`` for a +# sharded axis and a plain (rank-identical) tensor for a replicated one. +PlacementName = Literal["shard", "replicate"] +_VALID_PLACEMENTS: tuple[PlacementName, ...] = ("shard", "replicate") + + +# --------------------------------------------------------------------------- +# Configuration +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class DomainParallelConfig: + r"""Parsed ``domain_parallel`` configuration bound to its device mesh. + + Built once per reader by :meth:`from_dict`; every placement decision goes + through :meth:`decide`, so the raw dict is never re-read per sample. + + Parameters + ---------- + device_mesh : DeviceMesh + 1-D device mesh the batch axes are sharded over. + auto_shard_size : int + Auto gate: a batch axis shards when its length (tensor dim 0) is at + least this many entries and at least the world size. + placements : Mapping[str, PlacementName] + User overrides, dotted axis name -> ``"shard"`` | ``"replicate"``; + an entry applies to the axis it names and, by prefix, to every axis + beneath it. + """ + + device_mesh: DeviceMesh + auto_shard_size: int = DEFAULT_AUTO_SHARD_SIZE + placements: Mapping[str, PlacementName] = field(default_factory=dict) + + @classmethod + def from_dict( + cls, config: Mapping[str, Any] | None, device_mesh: DeviceMesh | None + ) -> DomainParallelConfig | None: + r"""Validate and parse the ``domain_parallel`` dict / ``device_mesh`` pair. + + Parameters + ---------- + config : Mapping or None + The ``domain_parallel`` configuration dict (see module docstring). + device_mesh : DeviceMesh or None + The device mesh the batch axes would be sharded over. + + Returns + ------- + DomainParallelConfig or None + ``None`` when both are absent (domain parallelism off). + + Raises + ------ + ValueError + On a missing/extra pairing, a non-1-D mesh, an unknown key, a + non-positive ``auto_shard_size``, or a bad placements entry. + """ + if config is None and device_mesh is None: + return None + if (config is None) != (device_mesh is None): + raise ValueError( + "domain_parallel and device_mesh must be provided together" + ) + if device_mesh.ndim != 1: + raise ValueError(f"device_mesh must be 1-D, got {device_mesh.ndim} dims") + + unknown = set(config) - {"auto_shard_size", "placements"} + if unknown: + raise ValueError( + f"unknown domain_parallel keys {sorted(unknown)}; " + 'expected "auto_shard_size" and/or "placements"' + ) + size = config.get("auto_shard_size", DEFAULT_AUTO_SHARD_SIZE) + if not isinstance(size, int) or isinstance(size, bool) or size < 1: + raise ValueError(f"auto_shard_size must be a positive int, got {size!r}") + placements = config.get("placements") or {} + if not isinstance(placements, Mapping): + raise ValueError( + f'placements must be a dict of axis -> "shard"|"replicate", ' + f"got {placements!r}" + ) + bad = {k: v for k, v in placements.items() if v not in _VALID_PLACEMENTS} + if bad: + raise ValueError(f'placements values must be "shard"|"replicate": {bad}') + return cls( + device_mesh=device_mesh, auto_shard_size=size, placements=dict(placements) + ) + + @property + def world_size(self) -> int: + r"""Number of ranks the batch axes are sharded over.""" + return self.device_mesh.size(0) + + def chunk_bounds(self, global_n: int) -> tuple[int, int]: + r"""This rank's ``[start, stop)`` share of an axis of length *global_n*.""" + return chunk_bounds(global_n, self.device_mesh) + + def placement_for( + self, axis: str, pinned: Mapping[str, PlacementName] | None = None + ) -> PlacementName | None: + r"""The override that applies to *axis*, if any. + + User ``placements`` win over reader *pinned* entries; within each, the + longest dotted prefix of *axis* wins. + + Parameters + ---------- + axis : str + Dotted axis name. + pinned : Mapping[str, PlacementName], optional + Reader-supplied structural pins (``"boundaries.stl": "replicate"``). + + Returns + ------- + PlacementName or None + ``"shard"`` / ``"replicate"`` if an override applies, else ``None``. + """ + overrides = {**(pinned or {}), **self.placements} + parts = axis.split(".") + for n in range(len(parts), 0, -1): + hit = overrides.get(".".join(parts[:n])) + if hit is not None: + return hit + return None + + def warn_unmatched(self, names: Iterable[str]) -> None: + r"""Log a warning for configured overrides that match none of *names*. + + Parameters + ---------- + names : Iterable[str] + Dotted names an override could legitimately target. + """ + names = list(names) + for key in self.placements: + if not any(n == key or n.startswith(key + ".") for n in names): + logger.warning( + "domain_parallel.placements[%r] matches no batch axis of this " + "sample (axes: %s); check the key", + key, + sorted(names), + ) + + def decide( + self, + axes: Mapping[str, int], + pinned: Mapping[str, PlacementName] | None = None, + *, + check_unmatched: bool = True, + ) -> dict[str, bool]: + r"""Decide shard-vs-replicate for each named batch axis. + + Parameters + ---------- + axes : Mapping[str, int] + Length (tensor dim 0) per axis name, from metadata; no data read. + pinned : Mapping[str, PlacementName], optional + Reader-supplied structural pins, overridden by ``placements``. + check_unmatched : bool, default=True + Warn about ``placements`` keys that match none of *axes*. + + Returns + ------- + dict[str, bool] + ``True`` to shard the axis, ``False`` to replicate it. Under the + gate an axis shards when its length is at least + ``auto_shard_size`` and at least the world size. + + Raises + ------ + ValueError + An axis pinned to ``"shard"`` with fewer entries than the world + size, which would leave a rank with an empty shard. + """ + if check_unmatched: + self.warn_unmatched(axes) + world_size = self.world_size + decisions: dict[str, bool] = {} + for axis, length in axes.items(): + pin = self.placement_for(axis, pinned) + if pin == "shard": + if length < world_size: + raise ValueError( + f"axis {axis!r} is pinned to shard but has {length} entries " + f"< world size {world_size}" + ) + decisions[axis] = True + elif pin == "replicate": + decisions[axis] = False + else: + decisions[axis] = length >= max(self.auto_shard_size, world_size) + return decisions + + +# --------------------------------------------------------------------------- +# Chunk arithmetic +# --------------------------------------------------------------------------- + + +def chunk_bounds(global_n: int, device_mesh: DeviceMesh) -> tuple[int, int]: + r"""This rank's ``[start, stop)`` share of an axis under ``torch.chunk`` semantics. + + Uses the same shard-shape arithmetic as + ``ShardTensor.from_local(sharding_shapes="chunk")``, so entries selected + with these bounds are exactly the local shard the later wrap declares. + + Parameters + ---------- + global_n : int + Global length of the batch axis being sharded. + device_mesh : DeviceMesh + 1-D device mesh the axis is sharded over. + + Returns + ------- + tuple[int, int] + Half-open ``[start, stop)`` range of entries owned by this rank. + """ + shapes = compute_sharding_shapes_from_chunking_global_shape( + device_mesh, PLACEMENTS, (global_n,) + ) + sizes = [s[0] for s in shapes[0]] + rank = device_mesh.get_local_rank(0) + start = sum(sizes[:rank]) + return start, start + sizes[rank] + + +# --------------------------------------------------------------------------- +# Placement resolution +# --------------------------------------------------------------------------- + + +def resolve_leaf_placements( + meta: dict[NestedKey, tuple[int, ...]], config: DomainParallelConfig +) -> dict[NestedKey, bool]: + r"""Decide shard-vs-replicate per leaf of a flat sample, by axis group. + + Leaves sharing a dim-0 length form one batch axis and are decided + together, so co-indexed arrays never end up with mismatched placements. + A ``placements`` entry naming any leaf pins its whole group; two leaves + of one group pinned differently is an error. Scalars always replicate. + + Parameters + ---------- + meta : dict[NestedKey, tuple[int, ...]] + Global shape per leaf (from store metadata; no data read). + config : DomainParallelConfig + Parsed configuration bound to the device mesh. + + Returns + ------- + dict[NestedKey, bool] + Per-leaf sharding decision. + """ + groups: dict[int, list[NestedKey]] = {} + for key, shape in meta.items(): + if len(shape) > 0: + groups.setdefault(shape[0], []).append(key) + + # One axis per dim-0 length, named by its leaves so errors and warnings + # can point at keys the user configured. Overrides apply by prefix. + axes: dict[str, int] = {} + pinned: dict[str, str] = {} + axis_of: dict[int, str] = {} + for length, keys in groups.items(): + names = sorted(key_to_str(k) for k in keys) + axis = "|".join(names) + axis_of[length] = axis + axes[axis] = length + pins = {} + for name in names: + pin = config.placement_for(name) + if pin is not None: + pins[name] = pin + if len(set(pins.values())) > 1: + raise ValueError( + f"leaves {sorted(pins)} share a batch axis (length {length}) but " + f"are pinned to different placements: {pins}" + ) + if pins: + pinned[axis] = next(iter(pins.values())) + + # Overrides were already applied per leaf above (they are pins now), so + # decide without the user placements; warn against the leaf names. + config.warn_unmatched(key_to_str(k) for k, shape in meta.items() if len(shape) > 0) + axis_decisions = replace(config, placements={}).decide( + axes, pinned, check_unmatched=False + ) + return { + key: (len(shape) > 0 and axis_decisions[axis_of[shape[0]]]) + for key, shape in meta.items() + } + + +# --------------------------------------------------------------------------- +# Host-stage payload +# --------------------------------------------------------------------------- + +# kind -> function rebuilding the final sample from the assembled TensorDict. +_REBUILDERS: dict[str, Callable[[TensorDict], Any]] = { + "tensordict": lambda td: td, +} + + +def register_proto_kind(kind: str, rebuild: Callable[[TensorDict], Any]) -> None: + r"""Register how a :class:`ShardedProto` of *kind* becomes a sample. + + Internal: the datapipes' own readers register ``"mesh"`` and + ``"domain_mesh"``; it is not a public extension point. + + Parameters + ---------- + kind : str + Payload kind tag (``"mesh"``, ``"domain_mesh"``, ...). + rebuild : Callable[[TensorDict], Any] + Builds the sample from the nested TensorDict whose sharded leaves + have already been wrapped as ShardTensors. + """ + _REBUILDERS[kind] = rebuild + + +@dataclass(frozen=True) +class ShardedProto: + r"""This rank's share of one sample, before ShardTensor assembly. + + The host-stage payload of a domain-parallel read. ``tensors`` is a + nested TensorDict mirroring the final sample's structure (for a mesh: + ``points``, ``cells``, ``point_data``, ``cell_data``, ``global_data``; + for a domain mesh: ``interior``, ``boundaries.``, ``global_data``). + Sharded leaves hold this rank's share; replicated leaves are complete. + + Readers return a proto instead of a finished sample; datasets move it to + the device (``to`` / ``pin_memory``, the same seam every sample flows + through) and then call :meth:`assemble`. + + Users aren't expected to interact with a ShardedProto, unless you're + assembling your own domain-parallel reader. + + Parameters + ---------- + tensors : TensorDict + Nested local tensors. + sharded : dict[tuple[str, ...], tuple[int, ...]] + Global shape per sharded leaf, keyed by nested tuple key. Leaves + absent from this map are replicated. + device_mesh : DeviceMesh + 1-D device mesh the selection was taken against; the wrap reuses it. + kind : str + Which registered rebuild turns the assembled TensorDict into the + sample (``"tensordict"``, ``"mesh"``, ``"domain_mesh"``). + """ + + tensors: TensorDict + sharded: dict[tuple[str, ...], tuple[int, ...]] + device_mesh: DeviceMesh + kind: str = "tensordict" + + def _replace_tensors(self, tensors: TensorDict) -> "ShardedProto": + return ShardedProto( + tensors=tensors, + sharded=self.sharded, + device_mesh=self.device_mesh, + kind=self.kind, + ) + + def to(self, device: torch.device, non_blocking: bool = False) -> "ShardedProto": + r"""Return a copy with the local tensors moved to *device*. + + Parameters + ---------- + device : torch.device + Target device. + non_blocking : bool, default=False + Passed through to ``TensorDict.to`` for async H2D copies. + """ + return self._replace_tensors(self.tensors.to(device, non_blocking=non_blocking)) + + def pin_memory(self) -> "ShardedProto": + r"""Return a copy with the local tensors in pinned host memory.""" + return self._replace_tensors(self.tensors.pin_memory()) + + def assemble(self) -> Any: + r"""Wrap sharded leaves as ``Shard(0)`` ShardTensors and rebuild the sample. + + The wrap is communication-free: every rank derives identical shard + shapes from the global shapes carried in :attr:`sharded`. + """ + try: + rebuild = _REBUILDERS[self.kind] + except KeyError: + raise ValueError( + f"no rebuild registered for proto kind {self.kind!r}; " + f"known kinds: {sorted(_REBUILDERS)}" + ) from None + return rebuild( + wrap_sharded_leaves(self.tensors, self.sharded, self.device_mesh) + ) + + +def assemble_if_proto(data: Any) -> Any: + r"""Assemble a :class:`ShardedProto` payload; pass anything else through.""" + return data.assemble() if isinstance(data, ShardedProto) else data + + +def wrap_sharded_leaves( + tensors: TensorDict, + sharded: dict[tuple[str, ...], tuple[int, ...]], + device_mesh: DeviceMesh, +) -> TensorDict: + r"""Rebuild *tensors* with every leaf in *sharded* wrapped as a ShardTensor. + + Structure (including empty sub-TensorDicts) is preserved; the result has + ``batch_size=[]`` at every level since sharded leaves carry global batch + lengths that the local sub-TensorDict batch sizes no longer match. The + wrap is the communication-free chunk path: every rank derives identical + shard shapes from the global shapes. + + Parameters + ---------- + tensors : TensorDict + Nested local tensors, already on the target device. + sharded : dict[tuple[str, ...], tuple[int, ...]] + Global shape per sharded leaf. + device_mesh : DeviceMesh + 1-D device mesh for the ``Shard(0)`` wrap. + + Returns + ------- + TensorDict + Same structure; sharded leaves are ``Shard(0)`` ShardTensors, + replicated leaves pass through as plain tensors. + """ + + def wrap(td: TensorDict, prefix: tuple[str, ...]) -> TensorDict: + out: dict[str, Any] = {} + for key, value in td.items(): + path = (*prefix, key) + if isinstance(value, TensorDict): + out[key] = wrap(value, path) + elif path in sharded: + out[key] = ShardTensor.from_local( + value, + device_mesh, + PLACEMENTS, + sharding_shapes="chunk", + global_shape=sharded[path], + ) + else: + out[key] = value + return TensorDict(out, batch_size=[]) + + return wrap(tensors, ()) + + +def as_leaf_key(key: NestedKey) -> tuple[str, ...]: + r"""Normalize a TensorDict leaf key to the tuple form :class:`ShardedProto` uses. + + A string is one component (it is a TensorDict key, not a dotted config + path); a tuple passes through. + """ + return (key,) if isinstance(key, str) else tuple(key) diff --git a/physicsnemo/datapipes/dataset.py b/physicsnemo/datapipes/dataset.py index 5548401de0..f5422a164f 100644 --- a/physicsnemo/datapipes/dataset.py +++ b/physicsnemo/datapipes/dataset.py @@ -186,6 +186,8 @@ def _load(self, index: int) -> tuple[TensorDict, dict[str, Any]]: if self.target_device is not None: data = data.to(self.target_device, non_blocking=True) + data = self._assemble(data) + if self.transforms is not None: data = self.transforms(data) @@ -349,6 +351,7 @@ def _consume( with preprocessing_stream(stream if use_stream else None): if self.target_device is not None: data = data.to(self.target_device, non_blocking=True) + data = self._assemble(data) if self.transforms is not None: data = self.transforms(data) diff --git a/physicsnemo/datapipes/mesh_dataset.py b/physicsnemo/datapipes/mesh_dataset.py index 789b231868..66365ab4d4 100644 --- a/physicsnemo/datapipes/mesh_dataset.py +++ b/physicsnemo/datapipes/mesh_dataset.py @@ -88,6 +88,10 @@ def __init__( ---------- reader : MeshReader or DomainMeshReader Mesh reader; returns (Mesh, metadata) or (DomainMesh, metadata). + A reader configured for domain-parallel reading (its + ``domain_parallel`` / ``device_mesh`` options) returns proto + payloads that this dataset assembles into ShardTensor-backed + meshes after the device transfer. transforms : sequence of MeshTransform, optional Transforms to apply in order. None means no transforms. device : str or torch.device, optional @@ -177,6 +181,8 @@ def _load( with torch.profiler.record_function("MeshDataset._load: data.to(device)"): data = data.to(self._device) + data = self._assemble(data) + for t in self.transforms: with torch.profiler.record_function( f"MeshDataset._load: transform {type(t).__name__}" @@ -292,6 +298,7 @@ def _apply_transforms(d: Any) -> Any: "MeshDataset._consume: data.to(device)" ): data = data.to(self._device, non_blocking=True) + data = self._assemble(data) with torch.profiler.record_function( "MeshDataset._consume: _apply_transforms" ): diff --git a/physicsnemo/datapipes/protocols.py b/physicsnemo/datapipes/protocols.py index 3d7e1d3b4e..4c633f48b8 100644 --- a/physicsnemo/datapipes/protocols.py +++ b/physicsnemo/datapipes/protocols.py @@ -44,6 +44,8 @@ import torch from tensordict import is_tensor_collection +from physicsnemo.datapipes._domain_parallel import assemble_if_proto + @contextlib.contextmanager def preprocessing_stream(stream: Optional["torch.cuda.Stream"]): @@ -118,10 +120,11 @@ class HostPayload: """A sample produced by the (thread-safe) I/O stage, staged on the host. A ``HostPayload`` is the boundary object between the I/O producer and - the main-thread consumer. It carries a CPU ``TensorDict`` (ideally - pinned, so the subsequent host-to-device copy can be asynchronous) - plus metadata. It is produced by a worker thread, which must not - launch device kernels. + the main-thread consumer. It carries a CPU sample -- a ``TensorDict``, + a mesh, or a domain-parallel ``ShardedProto`` of local rows -- ideally + pinned, so the subsequent host-to-device copy can be asynchronous, plus + metadata. It is produced by a worker thread, which must not launch + device kernels. Parameters ---------- @@ -130,7 +133,8 @@ class HostPayload: for map-style datasets, or any opaque descriptor for descriptor-driven sources. data : Any, optional - Host ``TensorDict`` (or mesh) payload. ``None`` on error. + Host ``TensorDict``, mesh, or ``ShardedProto`` payload. ``None`` on + error. metadata : dict, optional Per-sample metadata produced by the reader. error : Exception, optional @@ -234,6 +238,26 @@ def _pop_events(self) -> list: self._events_pending = [] return lst + @staticmethod + def _assemble(data: Any) -> Any: + """Assemble a domain-parallel proto payload; no-op for plain samples. + + Readers configured for domain-parallel reading return a + ``ShardedProto`` of this rank's rows; after the device transfer the + dataset turns it into the sample with ``Shard(0)`` ShardTensors. + + Parameters + ---------- + data : Any + Device-side payload. + + Returns + ------- + Any + The assembled sample, or *data* unchanged. + """ + return assemble_if_proto(data) + @abstractmethod def _load(self, index: int) -> tuple[Any, dict[str, Any]]: """Load and return a single sample ``(data, metadata)``. diff --git a/physicsnemo/datapipes/readers/__init__.py b/physicsnemo/datapipes/readers/__init__.py index 992c25829e..22741de985 100644 --- a/physicsnemo/datapipes/readers/__init__.py +++ b/physicsnemo/datapipes/readers/__init__.py @@ -22,8 +22,16 @@ - Converting to torch tensors - Async CPU->GPU transfers with optional prefetching - Returning Sample objects ready for the transform pipeline + +A custom reader can support domain-parallel (rank-local) reading by +implementing ``Reader._load_sample_domain_parallel``; ``DomainParallelConfig`` +and ``resolve_leaf_placements`` are the two helpers that implementation needs. """ +from physicsnemo.datapipes._domain_parallel import ( + DomainParallelConfig, + resolve_leaf_placements, +) from physicsnemo.datapipes.readers.base import Reader from physicsnemo.datapipes.readers.hdf5 import HDF5Reader from physicsnemo.datapipes.readers.mesh import DomainMeshReader, MeshReader @@ -34,6 +42,8 @@ __all__ = [ "Reader", + "DomainParallelConfig", + "resolve_leaf_placements", "HDF5Reader", "ZarrReader", "NumpyReader", diff --git a/physicsnemo/datapipes/readers/base.py b/physicsnemo/datapipes/readers/base.py index 2706eaa812..229b67a277 100644 --- a/physicsnemo/datapipes/readers/base.py +++ b/physicsnemo/datapipes/readers/base.py @@ -26,15 +26,24 @@ import logging from abc import ABC, abstractmethod -from typing import Any, Callable, Iterator +from typing import TYPE_CHECKING, Any, Callable, Iterator import numpy as np import torch from tensordict import TensorDict, is_leaf_nontensor +from physicsnemo.datapipes._domain_parallel import ( + DomainParallelConfig, + ShardedProto, + as_leaf_key, + resolve_leaf_placements, +) from physicsnemo.datapipes._indexing import _cyclic_block_indices from physicsnemo.datapipes._rng import spawn_generator +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh + logger = logging.getLogger(__name__) @@ -89,6 +98,8 @@ def __init__( pin_memory: bool = False, include_index_in_metadata: bool = True, coordinated_subsampling: dict[str, Any] | None = None, + domain_parallel: dict[str, Any] | None = None, + device_mesh: DeviceMesh | None = None, ) -> None: """ Initialize the reader. @@ -112,10 +123,51 @@ def __init__( point has equal inclusion probability while reads retain storage locality. This allows configuration via Hydra. Readers that don't support coordinated subsampling will ignore this parameter. + domain_parallel : dict[str, Any], optional + Optional dict to configure domain-parallel (rank-local) reading. + Requires ``device_mesh``. If provided, may contain: + + - ``auto_shard_size``: auto gate; a batch axis shards when its + length (the size of tensor dim 0) is at least this many + entries (default 1024), e.g. ``(200_000, 3)`` points shard + while a ``(512, 512)`` image replicates + - ``placements``: ``{key: "shard" | "replicate"}`` overrides; a + key pins the whole batch axis (every key sharing its dim-0 + length), applies by prefix to nested keys, and unnamed axes + fall back to the gate + + Supporting readers read only this rank's chunk of every sharded + key and return a proto payload the dataset assembles into + ``Shard(0)`` ShardTensors on the GPU. Composes with + ``coordinated_subsampling``: the rank chunk is taken of the + coordinated window, which requires a seed (``set_generator``) so + every rank draws the same window. Every rank in the device mesh + must request the same sample indices: use a sampler that hands + the same index sequence to every rank of a domain group (e.g. a + ``DistributedSampler`` over the data-parallel axis only). Readers + that don't support domain-parallel reading raise. + device_mesh : torch.distributed.device_mesh.DeviceMesh, optional + 1-D device mesh for domain-parallel reading. Not serializable + configuration: construct it in Python at runtime and inject it + alongside ``domain_parallel``. + + Raises + ------ + ValueError + If ``domain_parallel`` is given to a reader that does not support + domain-parallel reading, or the configuration is invalid. """ self.pin_memory = pin_memory self.include_index_in_metadata = include_index_in_metadata self._coordinated_subsampling_config = coordinated_subsampling + self._domain_parallel = DomainParallelConfig.from_dict( + domain_parallel, device_mesh + ) + if self._domain_parallel is not None and not self._supports_domain_parallel: + raise ValueError( + f"{type(self).__name__} does not support domain-parallel reading; " + "remove domain_parallel / device_mesh or use a supporting reader" + ) # Base seed + epoch for deterministic per-index RNG. See # :meth:`_index_generator`. ``None`` means no seed was provided # (random draws fall back to the global default RNG). @@ -215,6 +267,64 @@ def _supports_coordinated_subsampling(self) -> bool: """ return False + @property + def _supports_domain_parallel(self) -> bool: + """ + Return True if this reader supports domain-parallel reading. + + Override this property (and implement + :meth:`_load_sample_domain_parallel`) in subclasses that can read + rank-local chunks. + + Returns + ------- + bool + True if domain-parallel reading is supported. + """ + return False + + def _load_sample_domain_parallel( + self, index: int + ) -> tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]]: + """ + Load this rank's chunk of a single sample. + + This is the extension point for domain-parallel reading. Implement it + in a subclass that returns ``True`` from + :attr:`_supports_domain_parallel`; the base class validates the + configuration, wraps the result for the device transfer, and the + dataset assembles ``Shard(0)`` ShardTensors on the device. The + recipe is: + + 1. Collect the global shape of every key from metadata (no data read). + 2. ``selection, sharded = self._selection_plan(shapes, window, + windowed_keys, generator)`` decides shard-vs-replicate per key + (gate and ``placements``) and returns the dim-0 selection this + rank reads of each. + 3. Read ``array[selection[key]]`` per key and return the tensors with + ``sharded``. + + The same procedure with domain parallelism off reads whole arrays, so + one read method can back both ``_load_sample`` and this one. Must be + thread-safe (no scratch state on ``self``). If the reader subsamples, + the window must come from :meth:`_window_indices` with + :meth:`_index_generator`; a seed is then required so every rank draws + the same window. See ``ZarrReader`` for a complete implementation. + + Parameters + ---------- + index : int + Sample index (0 to len-1). + + Returns + ------- + tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]] + Per-key local tensors, and the global shape of every *sharded* + key. Keys absent from the map are replicated (complete on every + rank). + """ + raise NotImplementedError + @property def field_names(self) -> list[str]: """ @@ -255,7 +365,11 @@ def __getitem__(self, index: int) -> tuple[TensorDict, dict[str, Any]]: ) # Load data - data_dict = self._load_sample(index) + domain_parallel = self._domain_parallel is not None + if domain_parallel: + data_dict, sharded = self._load_sample_domain_parallel(index) + else: + data_dict = self._load_sample(index) # Build metadata metadata = self._get_sample_metadata(index) @@ -269,6 +383,14 @@ def __getitem__(self, index: int) -> tuple[TensorDict, dict[str, Any]]: if self.pin_memory: data = data.pin_memory() + if domain_parallel: + data = ShardedProto( + tensors=data, + sharded={as_leaf_key(k): tuple(v) for k, v in sharded.items()}, + device_mesh=self._domain_parallel.device_mesh, + kind="tensordict", + ) + return data, metadata def __iter__(self) -> Iterator[tuple[TensorDict, dict[str, Any]]]: @@ -351,6 +473,79 @@ def _index_generator(self, index: int) -> torch.Generator | None: return None return spawn_generator(self._seed_base, self._epoch, index) + def _selection_plan( + self, + shapes: dict[str, tuple[int, ...]], + window: np.ndarray | None, + windowed_keys: set[str], + generator: torch.Generator | None, + ) -> tuple[dict[str, slice | np.ndarray], dict[str, tuple[int, ...]]]: + """Decide which dim-0 selection of each array this rank reads. + + One procedure serves the plain and the domain-parallel read: with + domain parallelism off every key reads its window (or everything); + with it on, sharded keys read this rank's share of the window (or of + the full range) and replicated keys read as before. + + Parameters + ---------- + shapes : dict[str, tuple[int, ...]] + Effective global shape per array key: the window length on dim 0 + for windowed keys, the stored shape otherwise. Metadata only. + window : np.ndarray or None + Coordinated-subsampling window from :meth:`_window_indices`. + windowed_keys : set[str] + Keys the window applies to. + generator : torch.Generator or None + The per-sample generator; a seed is required when both + subsampling and domain parallelism are on. + + Returns + ------- + tuple[dict[str, slice | np.ndarray], dict[str, tuple[int, ...]]] + Selection to read per key, and the global shape of every sharded + key (empty when domain parallelism is off). + """ + selection: dict[str, slice | np.ndarray] = {} + sharded: dict[str, tuple[int, ...]] = {} + config = self._domain_parallel + if config is not None: + self._require_seed_for_domain_parallel(generator) + shard = resolve_leaf_placements(shapes, config) + for key, shape in shapes.items(): + windowed = window is not None and key in windowed_keys + if config is not None and shard[key]: + lo, hi = config.chunk_bounds(shape[0]) + # This rank's share of the window (a sub-slice of a 1-2 run + # cyclic block) or of the full range. + selection[key] = window[lo:hi] if windowed else slice(lo, hi) + sharded[key] = shape + else: + selection[key] = window if windowed else slice(None) + return selection, sharded + + def _require_seed_for_domain_parallel( + self, generator: torch.Generator | None + ) -> None: + """Domain-parallel subsampling needs a seed: every rank must draw the same window. + + Parameters + ---------- + generator : torch.Generator or None + The per-sample generator from :meth:`_index_generator`. + + Raises + ------ + ValueError + If coordinated subsampling is configured and no seed was set. + """ + if self._coordinated_subsampling_config is not None and generator is None: + raise ValueError( + "domain-parallel reading with coordinated_subsampling requires a " + "seed so every rank draws the same window: call set_generator on " + "the dataset (the DataLoader does this when given a seed)" + ) + def _window_indices( self, row_counts: Callable[[str], int | None], diff --git a/physicsnemo/datapipes/readers/mesh.py b/physicsnemo/datapipes/readers/mesh.py index 25869eff03..0712fe4f7f 100644 --- a/physicsnemo/datapipes/readers/mesh.py +++ b/physicsnemo/datapipes/readers/mesh.py @@ -26,16 +26,33 @@ import glob as _glob import logging +from dataclasses import dataclass from pathlib import Path -from typing import Any, Iterator +from typing import TYPE_CHECKING, Any, Iterator import torch - +from tensordict import TensorDict + +from physicsnemo.datapipes._domain_parallel import ( + DomainParallelConfig, + PlacementName, + ShardedProto, + as_leaf_key, + chunk_bounds, + register_proto_kind, +) from physicsnemo.datapipes._indexing import _cyclic_block_indices from physicsnemo.datapipes._rng import spawn_generator from physicsnemo.datapipes.registry import register from physicsnemo.mesh import DomainMesh, Mesh -from physicsnemo.mesh.calculus.measure import compose_measure_weights +from physicsnemo.mesh.calculus.measure import ( + MEASURE_WEIGHTS_KEY, + compose_measure_weights, +) +from physicsnemo.mesh.io import io_zarr + +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh logger = logging.getLogger(__name__) @@ -138,20 +155,25 @@ def _indices_to_runs(indices: torch.Tensor) -> list[tuple[int, int]]: def _zarr_mesh_subsampled( - group, + group: Any, n_cells: int | None, n_points: int | None, generator: torch.Generator | None, + *, + drop_cells: bool = False, ) -> Mesh: """Partial-read a zarr mesh group: fetch only the subsample window. Reproduces :func:`_subsample_mesh` semantics (cyclic contiguous blocks, vertex compaction, Horvitz-Thompson measure weights) while reading only the selected rows from the store instead of materializing the full mesh. + With ``drop_cells`` the group is read as a point cloud (cells and cell + data are never fetched), matching the reader's ``drop_interior_cells``. """ - from physicsnemo.mesh.io import io_zarr as _ioz - - total_cells = group["cells"].shape[0] if "cells" in group else 0 + _ioz = io_zarr + total_cells = ( + 0 if drop_cells else (group["cells"].shape[0] if "cells" in group else 0) + ) total_points = group["points"].shape[0] if total_cells > 0 and n_cells is not None and total_cells > n_cells: @@ -187,7 +209,18 @@ def _zarr_mesh_subsampled( point_data=_ioz._read_tree( group, "point_data", leaf_reader=lambda a: _ioz._read_rows(a, runs) ), - cell_data=_ioz._read_tree(group, "cell_data"), + cell_data=( + TensorDict({}, batch_size=[]) + if drop_cells + else _ioz._read_tree(group, "cell_data") + ), + global_data=_ioz._read_tree(group, "global_data"), + ) + + if drop_cells: + return Mesh( + points=_ioz._read_rows(group["points"], [(0, total_points)]), + point_data=_ioz._read_tree(group, "point_data"), global_data=_ioz._read_tree(group, "global_data"), ) @@ -214,8 +247,593 @@ def _subsample_mesh( return mesh +# --------------------------------------------------------------------------- +# Domain-parallel (rank-local) reading +# --------------------------------------------------------------------------- + + +# --------------------------------------------------------------------------- +# Rank-local (domain-parallel) reading +# --------------------------------------------------------------------------- +# +# A mesh is read one of two ways: a lazy memmap ``Mesh`` (``Mesh.load``) whose +# entries can be sliced without materializing the file, or a zarr group whose +# entries are fetched with ``io_zarr._read_rows`` / ``_read_index``. ``_read_mesh_selection`` +# hides that difference; everything above it plans in terms of dim-0 selections. +# +# Layout: one batch axis per mesh, always even chunks. +# - Point cloud (no cells): ``points`` / ``point_data`` chunk over the +# (windowed) point range; ``global_data`` replicates. +# - Mesh with cells: ``cells`` / ``cell_data`` chunk over the (windowed) cell +# range and ``points`` / ``point_data`` chunk over the point range -- after +# the same global compaction onto referenced vertices the eager reader +# performs when a cell window applies. Cells keep global vertex ids, so +# ``points[cells]`` on the assembled mesh is ShardTensor's routed gather. +# Point subsampling on a mesh with cells is not supported. +# Caches are dropped (they recompute lazily through ShardTensor ops). + +IndexSelection = slice | torch.Tensor +ShardedMap = dict[tuple[str, ...], tuple[int, ...]] +MeshSource = Mesh | Any # a lazy ``Mesh`` or a zarr group + + +def _is_zarr(src: MeshSource) -> bool: + return not isinstance(src, Mesh) + + +def _mesh_counts(src: MeshSource, drop_cells: bool = False) -> tuple[int, int]: + """``(n_points, n_cells)`` of a lazy mesh or zarr group; metadata only.""" + if _is_zarr(src): + n_points = src["points"].shape[0] + n_cells = src["cells"].shape[0] if "cells" in src else 0 + else: + n_points, n_cells = src.n_points, src.n_cells + return n_points, 0 if drop_cells else n_cells + + +def _read(t: torch.Tensor) -> torch.Tensor: + # clone() materializes memmap entries into plain host memory NOW, on the + # calling (worker) thread -- on a lazy mesh this is the actual disk read. + return torch.as_tensor(t).clone() + + +def _read_leaves(td: TensorDict, selection: IndexSelection | None = None) -> TensorDict: + """Read every (nested) leaf of a lazy TensorDict, optionally a dim-0 selection.""" + out = TensorDict({}, batch_size=[]) + for key, value in td.items(include_nested=True, leaves_only=True): + out.set(key, _read(value if selection is None else value[selection])) + return out + + +def _zarr_selection(arr: Any, selection: IndexSelection, n: int) -> torch.Tensor: + """Read a selection of a zarr array: a contiguous run, 1-2 runs, or scattered ids.""" + if isinstance(selection, slice): + start, stop, _ = selection.indices(n) + return io_zarr._read_rows(arr, [(start, stop)]) + runs = _indices_to_runs(selection) + if len(runs) <= 2: # a cyclic-block window: page-sequential run reads + return io_zarr._read_rows(arr, runs) + return io_zarr._read_index(arr, selection.numpy()) # scattered ids + + +def _read_cells(src: MeshSource, selection: IndexSelection) -> torch.Tensor: + """Connectivity only (used to compact a cell window before reading points).""" + if _is_zarr(src): + return _zarr_selection(src["cells"], selection, src["cells"].shape[0]) + return _read(src.cells[selection]) + + +def _read_mesh_selection( + src: MeshSource, + point_selection: IndexSelection, + cell_selection: IndexSelection, + *, + drop_cells: bool = False, +) -> Mesh: + """Read a selection of points and cells of a mesh into memory. + + Parameters + ---------- + src : Mesh or zarr group + A lazy memmap ``Mesh`` or an open zarr mesh group. + point_selection, cell_selection : slice or Tensor + Dim-0 selection of the point and cell axes (``point_data`` / + ``cell_data`` leaves follow their axis). + drop_cells : bool, default=False + Read the mesh as a point cloud: no cells or cell data. + + Returns + ------- + Mesh + Plain in-memory mesh holding exactly that selection; ``global_data`` + is read whole. Caches are not carried over. + """ + empty = TensorDict({}, batch_size=[]) + if _is_zarr(src): + n_points, n_cells = _mesh_counts(src) + points = _zarr_selection(src["points"], point_selection, n_points) + point_data = io_zarr._read_tree( + src, + "point_data", + leaf_reader=lambda a: _zarr_selection(a, point_selection, n_points), + ) + if drop_cells or n_cells == 0: + cells, cell_data = None, empty + else: + cells = _zarr_selection(src["cells"], cell_selection, n_cells) + cell_data = io_zarr._read_tree( + src, + "cell_data", + leaf_reader=lambda a: _zarr_selection(a, cell_selection, n_cells), + ) + global_data = io_zarr._read_tree(src, "global_data") + else: + points = _read(src.points[point_selection]) + point_data = _read_leaves(src.point_data, point_selection) + if drop_cells or src.n_cells == 0: + cells, cell_data = None, empty + else: + cells = _read(src.cells[cell_selection]) + cell_data = _read_leaves(src.cell_data, cell_selection) + global_data = _read_leaves(src.global_data) + return Mesh( + points=points, + cells=cells, + point_data=point_data, + cell_data=cell_data, + global_data=global_data, + ) + + +def _require_seed(generator: torch.Generator | None) -> None: + """Domain-parallel subsampling needs a seed: every rank must draw the same window. + + Without one the draw falls back to each process's global RNG, so ranks + would read different entries while agreeing on the global shapes -- a + silently corrupt sample. + """ + if generator is None: + raise ValueError( + "domain-parallel reading with subsampling requires a seed so every " + "rank draws the same window: call set_generator on the dataset (the " + "DataLoader does this when given a seed)" + ) + + +# ---- placement per sub-mesh ------------------------------------------------- + + +def mesh_axis_names(path: str) -> tuple[str, str]: + r"""``(points_axis, cells_axis)`` names for the sub-mesh at *path*. + + ``path`` is ``""`` for a standalone mesh, ``"interior"`` or + ``"boundaries."`` inside a domain mesh. A mesh has exactly one + batch axis: ``points`` for a point cloud, ``cells`` otherwise (its + vertices follow the cell partition). + """ + prefix = f"{path}." if path else "" + return f"{prefix}points", f"{prefix}cells" + + +def resolve_mesh_placements( + shapes: dict[str, tuple[int, int]], + config: DomainParallelConfig, + pinned: dict[str, PlacementName] | None = None, +) -> dict[str, tuple[bool, bool]]: + r"""Decide ``(shard_points, shard_cells)`` per sub-mesh from global counts. + + Parameters + ---------- + shapes : dict[str, tuple[int, int]] + ``path -> (n_points, n_cells)`` for every sub-mesh, where the counts + are the effective global counts (the subsample window length when + one applies). + config : DomainParallelConfig + Parsed configuration bound to the device mesh. + pinned : dict[str, PlacementName], optional + Reader-supplied structural pins (``"boundaries.stl": "replicate"``). + + Returns + ------- + dict[str, tuple[bool, bool]] + Per-path ``(shard_points, shard_cells)``. A point cloud is gated on + its ``points`` axis and never shards cells; a mesh with cells is + gated on its ``cells`` axis and its vertices follow that decision. + """ + axes: dict[str, int] = {} + for path, (n_points, n_cells) in shapes.items(): + points_axis, cells_axis = mesh_axis_names(path) + if n_cells > 0: + axes[cells_axis] = n_cells + else: + axes[points_axis] = n_points + decisions = config.decide(axes, pinned) + out: dict[str, tuple[bool, bool]] = {} + for path, (_n_points, n_cells) in shapes.items(): + points_axis, cells_axis = mesh_axis_names(path) + if n_cells > 0: + out[path] = (decisions[cells_axis], decisions[cells_axis]) + else: + out[path] = (decisions[points_axis], False) + return out + + +# ---- selection plan and read ------------------------------------------------ + + +@dataclass(frozen=True) +class MeshSelectionPlan: + r"""What this rank selects of one (sub-)mesh, decided from global counts only. + + Parameters + ---------- + shard_points, shard_cells : bool + Placement of the two batch axes. For a mesh with cells they are equal + (both follow the ``cells`` axis decision). + n_points_src, n_cells_src : int + Point and cell counts of the source mesh. + point_window : Tensor or None + Point-subsample cyclic block (point clouds only). The global point + count is then the window length. + cell_window : Tensor or None + Cell-subsample cyclic block (meshes with cells only). Selecting a + window compacts the point set to the vertices the window references + -- the same global operation the eager reader performs -- so every + rank reads the whole window of ``cells`` (integers, cheap) to derive + the identical referenced set, then reads only its chunk of it. + measure_factor : float or None + Horvitz-Thompson factor for a cell window (``n_cells_src / n_cells``). + device_mesh : DeviceMesh + 1-D mesh the chunks are taken against. + """ + + shard_points: bool + shard_cells: bool + n_points_src: int + n_cells_src: int + point_window: torch.Tensor | None + cell_window: torch.Tensor | None + measure_factor: float | None + device_mesh: DeviceMesh + + @property + def n_cells(self) -> int: + """Number of cells after windowing: the window length, else the source count.""" + return ( + len(self.cell_window) if self.cell_window is not None else self.n_cells_src + ) + + +def plan_mesh_selection( + n_points_src: int, + n_cells_src: int, + shard_points: bool, + shard_cells: bool, + device_mesh: DeviceMesh, + point_window: torch.Tensor | None = None, + cell_window: torch.Tensor | None = None, +) -> MeshSelectionPlan: + r"""Validate and bundle a :class:`MeshSelectionPlan`; see the class for the fields.""" + if n_cells_src > 0 and point_window is not None: + raise ValueError( + "point subsampling on a mesh with cells is not supported under " + "domain parallelism (it would remap connectivity globally); " + "subsample cells instead, or drop the cells" + ) + if n_cells_src == 0 and cell_window is not None: + raise ValueError("cell_window applies to meshes with cells only") + if n_cells_src > 0: + shard_points = shard_cells # vertices follow the cells axis + return MeshSelectionPlan( + shard_points=shard_points, + shard_cells=shard_cells, + n_points_src=n_points_src, + n_cells_src=n_cells_src, + point_window=point_window, + cell_window=cell_window, + measure_factor=( + n_cells_src / len(cell_window) if cell_window is not None else None + ), + device_mesh=device_mesh, + ) + + +def compact_cells(cells: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + r"""``(referenced_vertex_ids, remapped_cells)``: remap *cells* onto their vertices. + + ``referenced_vertex_ids`` is sorted, matching what ``Mesh.slice_cells`` + + ``slice_points`` produce; ``remapped_cells`` index into it. + """ + referenced, inverse = torch.unique(cells, return_inverse=True) + return referenced, inverse.reshape(cells.shape) + + +def _chunk(n: int, shard: bool, device_mesh: DeviceMesh) -> slice: + return slice(*chunk_bounds(n, device_mesh)) if shard else slice(None) + + +def _select(window: torch.Tensor | None, share: slice) -> IndexSelection: + return window[share] if window is not None else share + + +def mesh_selection_to_proto( + src: MeshSource, plan: MeshSelectionPlan, prefix: tuple[str, ...] = () +) -> tuple[TensorDict, ShardedMap]: + r"""Execute a :class:`MeshSelectionPlan` against a lazy mesh or zarr group. + + Three layouts, all even ``torch.chunk``-style splits so the assembly wrap + needs no communication: + + - point cloud: this rank's share of the (windowed) point range; + - mesh with cells, no window: this rank's share of the cell range and of + the point range -- cells keep their global vertex ids; + - mesh with cells and a cell window: the whole window of ``cells`` is + read (integers) and compacted onto its referenced vertices identically + on every rank; this rank then keeps its share of the remapped cells and + cell data and reads its share of the referenced vertices. + + Returns the nested proto TensorDict and its sharded-leaf map + (``path -> global shape``). Cells always index the global (compacted) + vertex space; ``points[cells]`` on the assembled mesh is the routed gather. + """ + device_mesh = plan.device_mesh + drop_cells = plan.n_cells_src == 0 + if drop_cells: + n_points = ( + len(plan.point_window) + if plan.point_window is not None + else plan.n_points_src + ) + p = _select(plan.point_window, _chunk(n_points, plan.shard_points, device_mesh)) + local = _read_mesh_selection(src, p, slice(0, 0), drop_cells=True) + cells = torch.zeros(0, 1, dtype=torch.long) + elif plan.cell_window is None: + n_points = plan.n_points_src + c = _chunk(plan.n_cells_src, plan.shard_cells, device_mesh) + p = _chunk(n_points, plan.shard_points, device_mesh) + local = _read_mesh_selection(src, p, c) + cells = local.cells + else: + window = plan.cell_window + referenced, remapped = compact_cells(_read_cells(src, window)) + n_points = len(referenced) + if plan.shard_points and n_points < device_mesh.size(0): + raise ValueError( + f"the cell window references only {n_points} vertices, fewer than " + f"the {device_mesh.size(0)} ranks sharding them; use a larger " + "subsample_n_cells or replicate this mesh" + ) + c = _chunk(len(window), plan.shard_cells, device_mesh) + p_ids = referenced[_chunk(n_points, plan.shard_points, device_mesh)] + local = _read_mesh_selection(src, p_ids, window[c]) + cells = remapped[c] + + cell_data = local.cell_data + if plan.measure_factor is not None: + weights = cell_data.get(MEASURE_WEIGHTS_KEY, None) + if weights is None: + weights = torch.ones(cells.shape[0], dtype=local.points.dtype) + cell_data[MEASURE_WEIGHTS_KEY] = weights * plan.measure_factor + + tensors = TensorDict( + { + "points": local.points, + "cells": cells, + "point_data": local.point_data, + "cell_data": cell_data, + "global_data": local.global_data, + }, + batch_size=[], + ) + + def leaves(td: TensorDict): + for key, value in td.items(include_nested=True, leaves_only=True): + yield as_leaf_key(key), value + + sharded: ShardedMap = {} + if plan.shard_points: + sharded[(*prefix, "points")] = (n_points, *local.points.shape[1:]) + for k, v in leaves(local.point_data): + sharded[(*prefix, "point_data", *k)] = (n_points, *v.shape[1:]) + if plan.shard_cells: + sharded[(*prefix, "cells")] = (plan.n_cells, *cells.shape[1:]) + for k, v in leaves(cell_data): + sharded[(*prefix, "cell_data", *k)] = (plan.n_cells, *v.shape[1:]) + return tensors, sharded + + +# ---- per sub-mesh preparation (phase A: metadata and windows only) ------------ + + +@dataclass +class _PreparedSubmesh: + """A sub-mesh ready for rank-local reading: source + effective global counts. + + ``src`` is a lazy memmap ``Mesh`` or an open zarr group. ``point_window`` / + ``cell_window`` are the subsample cyclic blocks (point clouds / meshes + with cells respectively); the effective global count of the windowed axis + is the window length. + """ + + src: MeshSource + n_points: int + n_cells: int + point_window: torch.Tensor | None = None + cell_window: torch.Tensor | None = None + + +def _prepare_submesh( + src: MeshSource, + *, + n_cells_sub: int | None, + n_points_sub: int | None, + generator: torch.Generator | None, + drop_cells: bool = False, +) -> _PreparedSubmesh: + """Phase A of a rank-local read: subsample decisions from metadata only. + + Nothing is read here. A mesh with cells draws its cell window; a point + cloud draws its point window. Point subsampling on a mesh with cells is + rejected (it remaps connectivity globally). The generator draw order + matches :func:`_subsample_mesh` (cells, then points), so every rank + derives the same windows. + """ + total_points, total_cells = _mesh_counts(src, drop_cells) + + if total_cells > 0: + if n_points_sub is not None: + raise NotImplementedError( + "subsample_n_points on a mesh with cells is not supported under " + "domain-parallel reading; use subsample_n_cells (or drop the " + "cells to read the mesh as a point cloud)" + ) + cell_window = None + if n_cells_sub is not None and total_cells > n_cells_sub: + _require_seed(generator) + cell_window = _cyclic_block_indices( + total_cells, n_cells_sub, generator=generator + ) + return _PreparedSubmesh( + src=src, + n_points=total_points, + n_cells=total_cells if cell_window is None else len(cell_window), + cell_window=cell_window, + ) + + point_window = None + if n_points_sub is not None and total_points > n_points_sub: + _require_seed(generator) + point_window = _cyclic_block_indices( + total_points, n_points_sub, generator=generator + ) + return _PreparedSubmesh( + src=src, + n_points=total_points if point_window is None else len(point_window), + n_cells=0, + point_window=point_window, + ) + + +def _read_prepared( + prep: _PreparedSubmesh, + shard_points: bool, + shard_cells: bool, + device_mesh: DeviceMesh, + prefix: tuple[str, ...] = (), +) -> tuple[TensorDict, ShardedMap]: + """Phase B: plan this rank's selection and read it into proto tensors.""" + plan = plan_mesh_selection( + *_mesh_counts(prep.src, drop_cells=prep.n_cells == 0), + shard_points, + shard_cells, + device_mesh, + point_window=prep.point_window, + cell_window=prep.cell_window, + ) + return mesh_selection_to_proto(prep.src, plan, prefix) + + +# ---- rebuild (device side, after the ShardTensor wrap) -------------------------- + + +def _rebuild_mesh(td: TensorDict) -> Mesh: + return Mesh( + points=td["points"], + cells=td["cells"], + point_data=td["point_data"], + cell_data=td["cell_data"], + global_data=td["global_data"], + ) + + +def _rebuild_domain_mesh(td: TensorDict) -> DomainMesh: + boundaries = td["boundaries"] + return DomainMesh( + interior=_rebuild_mesh(td["interior"]), + boundaries={ + name: _rebuild_mesh(boundaries[name]) for name in boundaries.keys() + }, + global_data=td["global_data"], + ) + + +register_proto_kind("mesh", _rebuild_mesh) +register_proto_kind("domain_mesh", _rebuild_domain_mesh) + + +class _MeshReaderBase: + """State and helpers shared by :class:`MeshReader` and :class:`DomainMeshReader`. + + Per-sample RNG (base seed + epoch), the cached zarr group handles, and + the domain-parallel configuration. + """ + + def _init_common( + self, domain_parallel: dict | None, device_mesh: DeviceMesh | None + ) -> None: + self._domain_parallel = DomainParallelConfig.from_dict( + domain_parallel, device_mesh + ) + # Base seed + epoch for deterministic per-index RNG (see + # :meth:`set_generator`). ``None`` means unseeded. + self._seed_base: int | None = None + self._epoch: int = 0 + self._zarr_groups: dict[Path, Any] = {} + + def _zarr_group(self, path: Path) -> Any: + """Open (and cache) the zarr group at *path*. + + Re-opening walks the store's group-metadata chain, and on networked + filesystems every uncached lookup is a metadata-server round-trip + per draw. + """ + group = self._zarr_groups.get(path) + if group is None: + group = self._zarr_groups[path] = io_zarr._open_group(path) + return group + + def _generator(self, index: int) -> torch.Generator | None: + """Per-sample generator from ``(base_seed, epoch, index)``; ``None`` if unseeded.""" + return ( + None + if self._seed_base is None + else spawn_generator(self._seed_base, self._epoch, index) + ) + + def set_generator(self, generator: torch.Generator) -> None: + """Assign a base seed for reproducible, order-independent subsampling. + + Called by :class:`MeshDataset` when the DataLoader provides a + seed. Stores ``generator.initial_seed()`` as the base seed; each + sample then derives its own generator from + ``(base_seed, epoch, index)``, so subsampling is reproducible + regardless of read order or worker thread. Required for + domain-parallel subsampling, where every rank must draw the same + window. + + Parameters + ---------- + generator : torch.Generator + Generator whose ``initial_seed()`` seeds all per-sample RNG. + """ + self._seed_base = generator.initial_seed() + + def set_epoch(self, epoch: int) -> None: + """Set the epoch used to vary per-sample RNG deterministically. + + The epoch is folded into each sample's derived seed, producing a + different (but deterministic) sequence of contiguous blocks each + epoch when a base seed has been assigned via :meth:`set_generator`. + """ + self._epoch = epoch + + def close(self) -> None: + """Release cached zarr store handles (``MeshDataset.close`` calls this).""" + self._zarr_groups.clear() + + @register() -class MeshReader: +class MeshReader(_MeshReaderBase): r""" Read single-mesh samples from directories of physicsnemo mesh files. @@ -232,6 +850,8 @@ def __init__( include_index_in_metadata: bool = True, subsample_n_points: int | None = None, subsample_n_cells: int | None = None, + domain_parallel: dict | None = None, + device_mesh: "torch.distributed.device_mesh.DeviceMesh | None" = None, ) -> None: """ Initialize the mesh reader. @@ -266,6 +886,22 @@ def __init__( probability as measure weights, preserving the integration measure (see :mod:`physicsnemo.mesh.calculus.measure`). Applied before ``subsample_n_points`` when both are set. + domain_parallel : dict, optional + Optional dict to configure domain-parallel (rank-local) + reading; see :mod:`physicsnemo.datapipes._domain_parallel` for + the schema. The mesh has two batch axes, ``points`` (points + + point_data) and ``cells`` (cell_data), each gated by + ``auto_shard_size`` or pinned via ``placements`` + (``{"points": "shard"}``). Composes with the ``subsample_*`` + options: a point cloud reads only this rank's chunk of the + (rank-consistent) subsample window; a mesh with cells is + subsampled in full first, since cell subsampling compacts the + point set globally. The sample returns as a proto payload that + ``MeshDataset`` assembles into a ShardTensor-backed ``Mesh`` on + the GPU. + device_mesh : torch.distributed.device_mesh.DeviceMesh, optional + 1-D device mesh for domain-parallel reading; required with + ``domain_parallel``. Constructed and injected at runtime. """ self._root = Path(path) self._pattern = pattern @@ -273,10 +909,7 @@ def __init__( self.include_index_in_metadata = include_index_in_metadata self.subsample_n_points = subsample_n_points self.subsample_n_cells = subsample_n_cells - # Base seed + epoch for deterministic per-index RNG (see - # :meth:`set_generator`). ``None`` means unseeded. - self._seed_base: int | None = None - self._epoch: int = 0 + self._init_common(domain_parallel, device_mesh) if not self._root.exists(): raise FileNotFoundError(f"Path not found: {self._root}") @@ -295,7 +928,7 @@ def _load_sample(self, index: int) -> Mesh: """Load a single Mesh from disk.""" mesh_path = self._paths[index] if (mesh_path / "zarr.json").exists(): - from physicsnemo.mesh.io import from_zarr, io_zarr + from physicsnemo.mesh.io import from_zarr if ( self.subsample_n_cells is not None @@ -305,29 +938,40 @@ def _load_sample(self, index: int) -> Mesh: # window from the store. The generator derivation matches # __getitem__, so the draw is identical to subsampling after # an eager load (whose subsample then no-ops). - # Cache opened store handles: re-opening walks the store's - # group-metadata chain, and on networked filesystems every - # uncached lookup is a metadata-server round-trip per draw. - cache = getattr(self, "_zarr_groups", None) - if cache is None: - cache = self._zarr_groups = {} - group = cache.get(mesh_path) - if group is None: - group = cache[mesh_path] = io_zarr._open_group(mesh_path) - generator = ( - None - if self._seed_base is None - else spawn_generator(self._seed_base, self._epoch, index) - ) return _zarr_mesh_subsampled( - group, + self._zarr_group(mesh_path), self.subsample_n_cells, self.subsample_n_points, - generator, + self._generator(index), ) return from_zarr(mesh_path) return Mesh.load(mesh_path) + def _load_domain_parallel(self, index: int) -> ShardedProto: + """Rank-local read: this rank's share of the (windowed) mesh as a proto.""" + mesh_path = self._paths[index] + sub_kw = dict( + n_cells_sub=self.subsample_n_cells, + n_points_sub=self.subsample_n_points, + generator=self._generator(index), + ) + if (mesh_path / "zarr.json").exists(): + prep = _prepare_submesh(self._zarr_group(mesh_path), **sub_kw) + else: + # Through ``_load_sample`` (a lazy memmap load) so subclass hooks + # that enrich the sample (e.g. merging external global_data) + # apply under domain-parallel reading too. + prep = _prepare_submesh(self._load_sample(index), **sub_kw) + + device_mesh = self._domain_parallel.device_mesh + (shard_points, shard_cells) = resolve_mesh_placements( + {"": (prep.n_points, prep.n_cells)}, self._domain_parallel + )[""] + tensors, sharded = _read_prepared(prep, shard_points, shard_cells, device_mesh) + return ShardedProto( + tensors=tensors, sharded=sharded, device_mesh=device_mesh, kind="mesh" + ) + def _get_sample_metadata(self, index: int) -> dict[str, Any]: """Return metadata for the sample (e.g. source path).""" return {"source_path": str(self._paths[index])} @@ -335,49 +979,25 @@ def _get_sample_metadata(self, index: int) -> dict[str, Any]: def __len__(self) -> int: return len(self._paths) - def set_generator(self, generator: torch.Generator) -> None: - """Assign a base seed for reproducible, order-independent subsampling. - - Called by :class:`MeshDataset` when the DataLoader provides a - seed. Stores ``generator.initial_seed()`` as the base seed; each - sample then derives its own generator from - ``(base_seed, epoch, index)``, so subsampling is reproducible - regardless of read order or worker thread. - - Parameters - ---------- - generator : torch.Generator - Generator whose ``initial_seed()`` seeds all per-sample RNG. - """ - self._seed_base = generator.initial_seed() - - def set_epoch(self, epoch: int) -> None: - """Set the epoch used to vary per-sample RNG deterministically. - - The epoch is folded into each sample's derived seed, producing a - different (but deterministic) sequence of contiguous blocks each - epoch when a base seed has been assigned via :meth:`set_generator`. - """ - self._epoch = epoch - def __getitem__(self, index: int) -> tuple[Mesh, dict[str, Any]]: metadata = self._get_sample_metadata(index) if self.include_index_in_metadata: metadata["index"] = index + if self._domain_parallel is not None: + # Rank-local read: only this rank's share leaves the store. The + # window is rank-consistent by the (seed, epoch, index) RNG + # scheme; the dataset assembles the ShardTensors on the GPU. + proto = self._load_domain_parallel(index) + return (proto.pin_memory() if self.pin_memory else proto), metadata + mesh = self._load_sample(index) - generator = ( - None - if self._seed_base is None - else spawn_generator(self._seed_base, self._epoch, index) - ) mesh = _subsample_mesh( mesh, self.subsample_n_cells, self.subsample_n_points, - generator=generator, + generator=self._generator(index), ) - if self.pin_memory: mesh = mesh.pin_memory() return mesh, metadata @@ -390,15 +1010,12 @@ def __iter__(self) -> Iterator[tuple[Mesh, dict[str, Any]]]: logger.error("Sample %s failed: %s", i, e) raise RuntimeError(f"Sample {i} failed: {e}") from e - def close(self) -> None: - """No reader-level resources; present for ``MeshDataset.close``.""" - def __repr__(self) -> str: return f"MeshReader(path={self._root!r}, len={len(self)})" @register() -class DomainMeshReader: +class DomainMeshReader(_MeshReaderBase): r""" Read DomainMesh samples from a directory of physicsnemo mesh files. @@ -419,6 +1036,8 @@ def __init__( extra_boundaries: dict[str, dict] | None = None, drop_interior_cells: bool = False, drop_in_file_boundaries: bool = False, + domain_parallel: dict | None = None, + device_mesh: "torch.distributed.device_mesh.DeviceMesh | None" = None, ) -> None: """ Initialize the domain mesh reader. @@ -492,6 +1111,18 @@ def __init__( otherwise be subsampled (an expensive ``slice_points`` remap, GIL-held, that blocks worker-thread overlap) and pinned every sample for nothing. + domain_parallel : dict, optional + Optional dict to configure domain-parallel (rank-local) + reading; see :class:`MeshReader`. Every sub-mesh's two batch + axes are gated independently and addressable in + ``placements`` by path: ``interior.points``, + ``boundaries..cells``, or ``boundaries.`` for both. + ``extra_boundaries`` are pinned to replicate by default (they + exist for whole-geometry queries such as SDF) but an explicit + ``placements`` entry overrides that. + device_mesh : torch.distributed.device_mesh.DeviceMesh, optional + 1-D device mesh for domain-parallel reading; required with + ``domain_parallel``. Constructed and injected at runtime. """ self._root = Path(path) self._pattern = pattern @@ -501,10 +1132,7 @@ def __init__( self.drop_in_file_boundaries = drop_in_file_boundaries self.subsample_n_points = subsample_n_points self.subsample_n_cells = subsample_n_cells - # Base seed + epoch for deterministic per-index RNG (see - # :meth:`set_generator`). ``None`` means unseeded. - self._seed_base: int | None = None - self._epoch: int = 0 + self._init_common(domain_parallel, device_mesh) self._extra_boundaries = extra_boundaries or {} if not self._root.exists(): @@ -524,7 +1152,7 @@ def _load_sample(self, index: int) -> DomainMesh: """Load a single DomainMesh from disk.""" path = self._paths[index] if (path / "zarr.json").exists(): - from physicsnemo.mesh.io import from_zarr, io_zarr + from physicsnemo.mesh.io import from_zarr if ( self.subsample_n_cells is not None @@ -534,22 +1162,14 @@ def _load_sample(self, index: int) -> DomainMesh: # sub-mesh); drop flags are honored at read time so skipped # data is never fetched. Generator derivation and sub-mesh # order match __getitem__, whose subsample then no-ops. - generator = ( - None - if self._seed_base is None - else spawn_generator(self._seed_base, self._epoch, index) - ) - cache = getattr(self, "_zarr_groups", None) - if cache is None: - cache = self._zarr_groups = {} - root = cache.get(path) - if root is None: - root = cache[path] = io_zarr._open_group(path) + generator = self._generator(index) + root = self._zarr_group(path) interior = _zarr_mesh_subsampled( root["interior"], self.subsample_n_cells, self.subsample_n_points, generator, + drop_cells=self.drop_interior_cells, ) boundaries = {} if not self.drop_in_file_boundaries and "boundaries" in root: @@ -573,32 +1193,105 @@ def _load_sample(self, index: int) -> DomainMesh: def __len__(self) -> int: return len(self._paths) - def set_generator(self, generator: torch.Generator) -> None: - """Assign a base seed for reproducible, order-independent subsampling. - - Called by :class:`MeshDataset` when the DataLoader provides a - seed. Stores ``generator.initial_seed()`` as the base seed; each - sample then derives its own generator from - ``(base_seed, epoch, index)``, so subsampling is reproducible - regardless of read order or worker thread. + def _load_domain_parallel(self, index: int) -> ShardedProto: + """Rank-local read of every sub-mesh; one proto for the whole domain. - Parameters - ---------- - generator : torch.Generator - Generator whose ``initial_seed()`` seeds all per-sample RNG. + Phase A prepares each sub-mesh (subsample decisions, window draws) + from metadata in the same order as the eager path draws its + generator; placements are then resolved jointly from the effective + global counts; phase B reads only this rank's share. Extra boundaries + are pinned to replicate unless ``placements`` says otherwise. """ - self._seed_base = generator.initial_seed() + path = self._paths[index] + generator = self._generator(index) + sub_kw = dict( + n_cells_sub=self.subsample_n_cells, + n_points_sub=self.subsample_n_points, + generator=generator, + ) - def set_epoch(self, epoch: int) -> None: - """Set the epoch used to vary per-sample RNG deterministically. + prepared: dict[str, _PreparedSubmesh] = {} + if (path / "zarr.json").exists(): + root = self._zarr_group(path) + prepared["interior"] = _prepare_submesh( + root["interior"], drop_cells=self.drop_interior_cells, **sub_kw + ) + if not self.drop_in_file_boundaries and "boundaries" in root: + for name, grp in root["boundaries"].groups(): + prepared[f"boundaries.{name}"] = _prepare_submesh(grp, **sub_kw) + global_data = io_zarr._read_tree(root, "global_data") + else: + # Through ``_load_sample`` (a lazy memmap load) so subclass hooks + # that enrich the sample apply under domain-parallel reading too. + domain = self._load_sample(index) + prepared["interior"] = _prepare_submesh( + domain.interior, drop_cells=self.drop_interior_cells, **sub_kw + ) + if not self.drop_in_file_boundaries: + for name in domain.boundary_names: + prepared[f"boundaries.{name}"] = _prepare_submesh( + domain.boundaries[name], **sub_kw + ) + global_data = domain.global_data + + # Extra boundaries: full resolution, never subsampled, replicate by + # default -- they exist for whole-geometry queries (e.g. SDF). + pinned: dict[str, PlacementName] = {} + for name, mesh in self._load_extra_boundary_meshes(index).items(): + prepared[f"boundaries.{name}"] = _PreparedSubmesh( + src=mesh, n_points=mesh.n_points, n_cells=mesh.n_cells + ) + pinned[f"boundaries.{name}"] = "replicate" - The epoch is folded into each sample's derived seed, producing a - different (but deterministic) sequence of contiguous blocks each - epoch when a base seed has been assigned via :meth:`set_generator`. - """ - self._epoch = epoch + device_mesh = self._domain_parallel.device_mesh + decisions = resolve_mesh_placements( + {p: (prep.n_points, prep.n_cells) for p, prep in prepared.items()}, + self._domain_parallel, + pinned=pinned, + ) + + sharded: ShardedMap = {} + boundaries: dict[str, TensorDict] = {} + interior = None + for p, prep in prepared.items(): + shard_points, shard_cells = decisions[p] + prefix = tuple(p.split(".")) + td, sub = _read_prepared( + prep, shard_points, shard_cells, device_mesh, prefix + ) + sharded.update(sub) + if p == "interior": + interior = td + else: + boundaries[prefix[1]] = td + + tensors = TensorDict( + { + "interior": interior, + "boundaries": TensorDict(boundaries, batch_size=[]), + # Nesting-aware: global_data may hold sub-TensorDicts. + "global_data": _read_leaves(global_data), + }, + batch_size=[], + ) + return ShardedProto( + tensors=tensors, + sharded=sharded, + device_mesh=device_mesh, + kind="domain_mesh", + ) def __getitem__(self, index: int) -> tuple[DomainMesh, dict[str, Any]]: + if self._domain_parallel is not None: + proto = self._load_domain_parallel(index) + metadata: dict[str, Any] = { + "source_path": str(self._paths[index]), + "boundary_names": sorted(proto.tensors["boundaries"].keys()), + } + if self.include_index_in_metadata: + metadata["index"] = index + return (proto.pin_memory() if self.pin_memory else proto), metadata + dm = self._load_sample(index) # Trim unused data before subsample/pin. Both references are lazy (no @@ -627,11 +1320,7 @@ def __getitem__(self, index: int) -> tuple[DomainMesh, dict[str, Any]]: ) if self.subsample_n_cells is not None or self.subsample_n_points is not None: - generator = ( - None - if self._seed_base is None - else spawn_generator(self._seed_base, self._epoch, index) - ) + generator = self._generator(index) sub_kw = dict( n_cells=self.subsample_n_cells, n_points=self.subsample_n_points, @@ -661,6 +1350,7 @@ def __getitem__(self, index: int) -> tuple[DomainMesh, dict[str, Any]]: if self.pin_memory: dm = dm.pin_memory() + return dm, metadata def _load_extra_boundaries(self, dm: DomainMesh, index: int) -> DomainMesh: @@ -717,8 +1407,5 @@ def __iter__(self) -> Iterator[tuple[DomainMesh, dict[str, Any]]]: logger.error("Sample %s failed: %s", i, e) raise RuntimeError(f"Sample {i} failed: {e}") from e - def close(self) -> None: - """No reader-level resources; present for ``MeshDataset.close``.""" - def __repr__(self) -> str: return f"DomainMeshReader(path={self._root!r}, len={len(self)})" diff --git a/physicsnemo/datapipes/readers/tensorstore_zarr.py b/physicsnemo/datapipes/readers/tensorstore_zarr.py index ece6dd603e..689ea31f7e 100644 --- a/physicsnemo/datapipes/readers/tensorstore_zarr.py +++ b/physicsnemo/datapipes/readers/tensorstore_zarr.py @@ -26,7 +26,7 @@ import importlib import json from pathlib import Path -from typing import Any, Optional +from typing import TYPE_CHECKING, Any, Optional import torch @@ -34,6 +34,9 @@ from physicsnemo.datapipes.readers.base import Reader from physicsnemo.datapipes.registry import register +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh + # Check if tensorstore is available TENSORSTORE_AVAILABLE = check_version_spec("tensorstore", hard_fail=False) @@ -97,6 +100,8 @@ def __init__( pin_memory: bool = False, include_index_in_metadata: bool = True, coordinated_subsampling: Optional[dict[str, Any]] = None, + domain_parallel: Optional[dict[str, Any]] = None, + device_mesh: Optional["DeviceMesh"] = None, ) -> None: """ Initialize the TensorStore Zarr reader. @@ -128,6 +133,14 @@ def __init__( coordinated_subsampling : dict[str, Any], optional Optional dict to configure coordinated subsampling. If provided, must contain ``n_points`` (int) and ``target_keys`` (list of str). + domain_parallel : dict[str, Any], optional + Optional dict to configure domain-parallel (rank-local) reading; + see :class:`~physicsnemo.datapipes.readers.base.Reader`. Sharded + keys read only this rank's chunk (of the coordinated window when + subsampling is also configured). + device_mesh : torch.distributed.device_mesh.DeviceMesh, optional + 1-D device mesh for domain-parallel reading; required with + ``domain_parallel``. Constructed and injected at runtime. Raises ------ @@ -149,6 +162,8 @@ def __init__( pin_memory=pin_memory, include_index_in_metadata=include_index_in_metadata, coordinated_subsampling=coordinated_subsampling, + domain_parallel=domain_parallel, + device_mesh=device_mesh, ) self.path = Path(path).expanduser().resolve() @@ -315,26 +330,52 @@ def _finalize_sample( data[key] = default_value.clone() return data - def _load_sample(self, index: int) -> dict[str, torch.Tensor]: - """Load a single sample from a Zarr group using TensorStore.""" + def _read_sample( + self, index: int + ) -> tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]]: + """Read a sample: the whole window, or this rank's share of it. + + Store opens are metadata-only, so placements (when domain parallelism + is on) resolve before any data is read. Attributes and default values + always replicate. + """ # Per-sample generator: reproducible regardless of read order/thread. generator = self._index_generator(index) stores, attributes = self._open_stores(index) - subsample_indices, target_keys_set = self._window_indices( + window, windowed_keys = self._window_indices( lambda key: stores[key].shape[0] if key in stores else None, generator, ) - # Trigger async reads - tensor_futures = {} - for key, store in stores.items(): - # Apply subsampling if this key is a target - if subsample_indices is not None and key in target_keys_set: - tensor_futures[key] = store[subsample_indices].read() - else: - tensor_futures[key] = store[:].read() + shapes = { + key: ( + len(window) + if window is not None and key in windowed_keys + else store.shape[0], + *store.shape[1:], + ) + for key, store in stores.items() + } + selection, sharded = self._selection_plan( + shapes, window, windowed_keys, generator + ) - return self._finalize_sample(tensor_futures, attributes) + # Trigger async reads of each key's selection + tensor_futures = { + key: store[selection[key]].read() for key, store in stores.items() + } + return self._finalize_sample(tensor_futures, attributes), sharded + + def _load_sample(self, index: int) -> dict[str, torch.Tensor]: + """Load a single sample from a Zarr group using TensorStore.""" + data, _ = self._read_sample(index) + return data + + def _load_sample_domain_parallel( + self, index: int + ) -> tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]]: + """Load this rank's chunk of a single sample using TensorStore.""" + return self._read_sample(index) def __len__(self) -> int: """Return number of samples.""" @@ -356,6 +397,11 @@ def _supports_coordinated_subsampling(self) -> bool: """TensorStore Zarr reader supports coordinated subsampling.""" return True + @property + def _supports_domain_parallel(self) -> bool: + """TensorStore Zarr reader supports domain-parallel reading.""" + return True + def __repr__(self) -> str: subsample_info = "" if self._coordinated_subsampling_config is not None: diff --git a/physicsnemo/datapipes/readers/zarr.py b/physicsnemo/datapipes/readers/zarr.py index f6baa9d095..53bd165aef 100644 --- a/physicsnemo/datapipes/readers/zarr.py +++ b/physicsnemo/datapipes/readers/zarr.py @@ -23,7 +23,7 @@ from __future__ import annotations from pathlib import Path -from typing import Any, Optional +from typing import TYPE_CHECKING, Any, Optional import numpy as np import torch @@ -32,6 +32,9 @@ from physicsnemo.datapipes.readers.base import Reader from physicsnemo.datapipes.registry import register +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh + zarr = OptionalImport("zarr") @@ -87,6 +90,8 @@ def __init__( pin_memory: bool = False, include_index_in_metadata: bool = True, coordinated_subsampling: Optional[dict[str, Any]] = None, + domain_parallel: Optional[dict[str, Any]] = None, + device_mesh: Optional["DeviceMesh"] = None, cache_stores: bool = True, ) -> None: """ @@ -115,6 +120,14 @@ def __init__( coordinated_subsampling : dict[str, Any], optional Optional dict to configure coordinated subsampling. If provided, must contain ``n_points`` (int) and ``target_keys`` (list of str). + domain_parallel : dict[str, Any], optional + Optional dict to configure domain-parallel (rank-local) reading; + see :class:`~physicsnemo.datapipes.readers.base.Reader`. Sharded + keys read only this rank's chunk (of the coordinated window when + subsampling is also configured). + device_mesh : torch.distributed.device_mesh.DeviceMesh, optional + 1-D device mesh for domain-parallel reading; required with + ``domain_parallel``. Constructed and injected at runtime. cache_stores : bool, default=True If True, cache opened zarr stores to avoid repeated opening and prevent executor shutdown errors. Set to False if memory is a @@ -136,6 +149,8 @@ def __init__( pin_memory=pin_memory, include_index_in_metadata=include_index_in_metadata, coordinated_subsampling=coordinated_subsampling, + domain_parallel=domain_parallel, + device_mesh=device_mesh, ) self.path = Path(path).expanduser().resolve() @@ -272,46 +287,68 @@ def _array_rows(self, array) -> int: """Row count of an array's batch dim (dim 1 in single-group mode).""" return array.shape[1] if self._single_group_mode else array.shape[0] - def _load_sample(self, index: int) -> dict[str, torch.Tensor]: - """Load a single sample from a Zarr group.""" + def _read_sample( + self, index: int + ) -> tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]]: + """Read a sample: the whole window, or this rank's share of it. + + Placements (when domain parallelism is on) resolve from store + metadata shapes before any data is read. Attributes and default + values always replicate. + """ # Per-sample generator: reproducible regardless of read order/thread. generator = self._index_generator(index) root, available_attrs = self._open_sample(index) - subsample_indices, target_keys_set = self._window_indices( + window, windowed_keys = self._window_indices( lambda key: self._array_rows(root[key]) if key in root else None, generator, ) - data = {} - # Load each field + # Effective global shape per array key (window length on dim 0 for + # windowed keys); metadata only. + shapes: dict[str, tuple[int, ...]] = {} + for field in self.fields: + if field in root: + array = root[field] + n_rows = ( + len(window) + if window is not None and field in windowed_keys + else self._array_rows(array) + ) + trailing = ( + array.shape[2:] if self._single_group_mode else array.shape[1:] + ) + shapes[field] = (n_rows, *trailing) + selection, sharded = self._selection_plan( + shapes, window, windowed_keys, generator + ) + + data: dict[str, torch.Tensor] = {} for field in self.fields: if field in root: if self._single_group_mode: - # Single group mode: index into first dimension - if subsample_indices is not None and field in target_keys_set: - # Apply subsampling on dimensions after the first - data[field] = torch.from_numpy( - root[field][index, subsample_indices] - ) - else: - data[field] = torch.from_numpy(root[field][index]) + # Single group mode: the sample is dim 0, the batch axis is dim 1. + data[field] = torch.from_numpy(root[field][index, selection[field]]) else: - # Directory mode: load entire array or subsample - if subsample_indices is not None and field in target_keys_set: - data[field] = torch.from_numpy(root[field][subsample_indices]) - else: - data[field] = torch.from_numpy(root[field][:]) - + data[field] = torch.from_numpy(root[field][selection[field]]) elif field in available_attrs: # Load from attributes (discovered at runtime for this sample) - attr_value = root.attrs[field] - data[field] = self._convert_attr_to_tensor(attr_value, field) - + data[field] = self._convert_attr_to_tensor(root.attrs[field], field) elif field in self.default_values: data[field] = self.default_values[field].clone() + return data, sharded + def _load_sample(self, index: int) -> dict[str, torch.Tensor]: + """Load a single sample from a Zarr group.""" + data, _ = self._read_sample(index) return data + def _load_sample_domain_parallel( + self, index: int + ) -> tuple[dict[str, torch.Tensor], dict[str, tuple[int, ...]]]: + """Load this rank's chunk of a single sample from a Zarr group.""" + return self._read_sample(index) + def _convert_attr_to_tensor(self, value: Any, field_name: str) -> torch.Tensor: """ Convert an attribute value to a torch.Tensor. @@ -382,6 +419,11 @@ def _supports_coordinated_subsampling(self) -> bool: """Zarr reader supports coordinated subsampling.""" return True + @property + def _supports_domain_parallel(self) -> bool: + """Zarr reader supports domain-parallel (rank-local) reading.""" + return True + def close(self) -> None: """Close resources and cached zarr stores.""" # Clear cached stores to allow garbage collection diff --git a/physicsnemo/datapipes/transforms/mesh/transforms.py b/physicsnemo/datapipes/transforms/mesh/transforms.py index 73bdb8bb33..b58502dcc6 100644 --- a/physicsnemo/datapipes/transforms/mesh/transforms.py +++ b/physicsnemo/datapipes/transforms/mesh/transforms.py @@ -212,8 +212,16 @@ def _compute_com(self, mesh: Mesh) -> Float[torch.Tensor, " spatial_dims"]: areas = mesh.cell_areas # (n_cells,) centroids = mesh.cell_centroids # (n_cells, n_spatial_dims) total_area = areas.sum() - return (centroids * areas.unsqueeze(-1)).sum(dim=0) / total_area - return mesh.points.mean(dim=0) + com = (centroids * areas.unsqueeze(-1)).sum(dim=0) / total_area + else: + com = mesh.points.mean(dim=0) + # On a sharded mesh the reduction yields a ShardTensor. Materialize + # it: the offset broadcasts against plain and sharded tensors alike, + # whereas a ShardTensor offset would promote every plain sub-mesh it + # touches (e.g. a replicated SDF reference surface) into ShardTensors. + if hasattr(com, "full_tensor"): + com = com.full_tensor() + return com def __call__(self, mesh: Mesh) -> Mesh: return mesh.translate(-self._compute_com(mesh)) diff --git a/test/datapipes/test_domain_parallel_config.py b/test/datapipes/test_domain_parallel_config.py new file mode 100644 index 0000000000..530bd9b92a --- /dev/null +++ b/test/datapipes/test_domain_parallel_config.py @@ -0,0 +1,374 @@ +# 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"""Pure-logic tests for the domain-parallel reader configuration. + +Placement resolution reads only global metadata and the mesh size, so it is +testable single-rank with a stub device mesh; the distributed read/wrap +behavior is covered in ``test/domain_parallel/datapipes/``. +""" + +import dataclasses +import logging + +import pytest +import torch +from tensordict import TensorDict + +from physicsnemo.datapipes._domain_parallel import ( + DomainParallelConfig, + ShardedProto, + assemble_if_proto, + resolve_leaf_placements, +) +from physicsnemo.datapipes.readers.base import Reader +from physicsnemo.datapipes.readers.mesh import ( + mesh_selection_to_proto, + plan_mesh_selection, + resolve_mesh_placements, +) +from physicsnemo.mesh import Mesh +from physicsnemo.mesh.calculus.measure import MEASURE_WEIGHTS_KEY + + +class _StubMesh: + """Duck-typed 1-D device mesh: just a world size.""" + + def __init__(self, world_size: int): + self._world_size = world_size + + @property + def ndim(self) -> int: + return 1 + + def size(self, dim: int = 0) -> int: + return self._world_size + + +def test_validate_requires_pairing(): + DomainParallelConfig.from_dict(None, None) # both absent: fine + with pytest.raises(ValueError, match="together"): + DomainParallelConfig.from_dict({}, None) + with pytest.raises(ValueError, match="together"): + DomainParallelConfig.from_dict(None, _StubMesh(2)) + + +def test_validate_rejects_bad_config(): + mesh = _StubMesh(2) + with pytest.raises(ValueError, match="unknown"): + DomainParallelConfig.from_dict({"nonsense": {}}, mesh) + with pytest.raises(ValueError, match="auto_shard_size"): + DomainParallelConfig.from_dict({"auto_shard_size": 0}, mesh) + with pytest.raises(ValueError, match="placements must be a dict"): + DomainParallelConfig.from_dict({"placements": "auto"}, mesh) + with pytest.raises(ValueError, match="shard"): + DomainParallelConfig.from_dict({"placements": {"x": "banana"}}, mesh) + DomainParallelConfig.from_dict( + {"auto_shard_size": 8, "placements": {"x": "shard"}}, mesh + ) + DomainParallelConfig.from_dict({}, mesh) + + +@pytest.mark.parametrize( + ("rows", "config", "world", "expected"), + [ + # Default gate: dim 0 at least 1024 long, whatever the world size. + (1024, {}, 4, True), + (1023, {}, 4, False), + (200_000, {}, 2, True), + # Explicit gate. + (64, {"auto_shard_size": 32}, 2, True), + (31, {"auto_shard_size": 32}, 2, False), + # Never shorter than the world size under the gate. + (3, {"auto_shard_size": 1}, 4, False), + (4, {"auto_shard_size": 1}, 4, True), + # Pins bypass the gate. + (3, {"auto_shard_size": 10**6, "placements": {"a": "shard"}}, 2, True), + (10**9, {"placements": {"a": "replicate"}}, 2, False), + ], +) +def test_axis_gate(rows, config, world, expected): + assert DomainParallelConfig.from_dict(config, _StubMesh(world)).decide( + {"a": rows} + ) == {"a": expected} + + +def test_axis_pin_below_world_size_raises(): + with pytest.raises(ValueError, match="world size"): + DomainParallelConfig.from_dict( + {"placements": {"a": "shard"}}, _StubMesh(4) + ).decide({"a": 3}) + + +def test_axis_pin_by_prefix_and_reader_pins(): + axes = { + "interior.points": 10**6, + "interior.cells": 10**6, + "boundaries.stl.points": 10**6, + "boundaries.stl.cells": 10**6, + "boundaries.wing.points": 10**6, + } + mesh = _StubMesh(2) + # Reader pin on the sub-mesh applies to both of its axes. + pinned = {"boundaries.stl": "replicate"} + decisions = DomainParallelConfig.from_dict({}, mesh).decide(axes, pinned) + assert decisions["boundaries.stl.points"] is False + assert decisions["boundaries.stl.cells"] is False + assert decisions["interior.points"] is True + assert decisions["boundaries.wing.points"] is True + # User config overrides the reader pin; a deeper key beats a shallower one. + config = { + "placements": { + "boundaries.stl": "shard", + "interior": "replicate", + "interior.points": "shard", + } + } + decisions = DomainParallelConfig.from_dict(config, mesh).decide(axes, pinned) + assert decisions["boundaries.stl.points"] is True + assert decisions["interior.points"] is True + assert decisions["interior.cells"] is False + + +def test_leaf_placements_group_by_axis(): + mesh = _StubMesh(4) + meta = { + "coords": (5000, 3), + "fields": (5000, 4), + "mask": (5000,), + "params": (7,), + "scalar": (), + } + # One decision per shared dim-0 length; scalars always replicate. + decisions = resolve_leaf_placements(meta, DomainParallelConfig.from_dict({}, mesh)) + assert decisions == { + "coords": True, + "fields": True, + "mask": True, + "params": False, + "scalar": False, + } + # Pinning one leaf pins its whole group. + decisions = resolve_leaf_placements( + meta, + DomainParallelConfig.from_dict({"placements": {"mask": "replicate"}}, mesh), + ) + assert not any(decisions[k] for k in ("coords", "fields", "mask")) + # Conflicting pins within a group are an error. + with pytest.raises(ValueError, match="different placements"): + resolve_leaf_placements( + meta, + DomainParallelConfig.from_dict( + {"placements": {"coords": "shard", "fields": "replicate"}}, mesh + ), + ) + + +def test_mesh_placements_axes(): + mesh = _StubMesh(2) + shapes = {"interior": (10**5, 0), "boundaries.wing": (10**5, 5 * 10**4)} + decisions = resolve_mesh_placements( + shapes, DomainParallelConfig.from_dict({}, mesh) + ) + # A point cloud is gated on points and never shards cells; a mesh with + # cells is gated on cells and its vertices follow. + assert decisions["interior"] == (True, False) + assert decisions["boundaries.wing"] == (True, True) + decisions = resolve_mesh_placements( + shapes, + DomainParallelConfig.from_dict( + {"placements": {"boundaries.wing.cells": "replicate"}}, mesh + ), + ) + assert decisions["boundaries.wing"] == (False, False) + # A few cells but many points: the cells axis decides, so it replicates. + decisions = resolve_mesh_placements( + {"": (10**6, 3)}, DomainParallelConfig.from_dict({}, mesh) + ) + assert decisions[""] == (False, False) + + +class _StubDeviceMesh(_StubMesh): + """Stub with a rank, enough for ``chunk_bounds``.""" + + def __init__(self, world_size: int, rank: int): + super().__init__(world_size) + self._rank = rank + + def get_local_rank(self, dim: int = 0) -> int: + return self._rank + + def get_coordinate(self): + return [self._rank] + + +def _small_mesh() -> Mesh: + torch.manual_seed(0) + return Mesh( + points=torch.randn(10, 3), + cells=torch.tensor([[0, 1, 2], [3, 4, 5], [6, 7, 8], [2, 3, 9]]), + point_data={"t": torch.randn(10)}, + cell_data={"p": torch.randn(4, 2)}, + global_data={"Re": torch.tensor(1.0)}, + ) + + +def test_mesh_selection_to_proto_point_cloud_window(): + """A point cloud reads its chunk of the (windowed) point range and reports + the window length as the global count.""" + mesh = _small_mesh() + cloud = Mesh(points=mesh.points, point_data={"t": mesh.point_data["t"]}) + dm = _StubDeviceMesh(2, 1) + plan = plan_mesh_selection(10, 0, True, False, dm) + tensors, sharded = mesh_selection_to_proto(cloud, plan, ("interior",)) + torch.testing.assert_close(tensors["points"], mesh.points[5:10]) + torch.testing.assert_close(tensors["point_data", "t"], mesh.point_data["t"][5:10]) + assert sharded == { + ("interior", "points"): (10, 3), + ("interior", "point_data", "t"): (10,), + } + + window = torch.tensor([7, 8, 9, 0, 1, 2]) + plan = plan_mesh_selection(10, 0, True, False, dm, point_window=window) + tensors, sharded = mesh_selection_to_proto(cloud, plan) + torch.testing.assert_close(tensors["points"], mesh.points[[0, 1, 2]]) + assert sharded[("points",)] == (6, 3) + + +def test_mesh_selection_to_proto_cells_no_window(): + """Full resolution: chunk of the cell range and chunk of the point range; + cells keep their global vertex ids.""" + mesh = _small_mesh() + dm = _StubDeviceMesh(2, 0) + plan = plan_mesh_selection(10, 4, True, True, dm) + tensors, sharded = mesh_selection_to_proto(mesh, plan, ("b",)) + torch.testing.assert_close(tensors["points"], mesh.points[:5]) + torch.testing.assert_close(tensors["cells"], mesh.cells[:2]) + torch.testing.assert_close(tensors["cell_data", "p"], mesh.cell_data["p"][:2]) + assert sharded == { + ("b", "points"): (10, 3), + ("b", "point_data", "t"): (10,), + ("b", "cells"): (4, 3), + ("b", "cell_data", "p"): (4, 2), + } + # Replicated: everything, untouched. + plan = plan_mesh_selection(10, 4, False, False, dm) + tensors, sharded = mesh_selection_to_proto(mesh, plan) + torch.testing.assert_close(tensors["points"], mesh.points) + torch.testing.assert_close(tensors["cells"], mesh.cells) + assert sharded == {} + + +def test_mesh_selection_to_proto_cell_window_compacts_globally(): + """A cell window compacts onto the referenced vertices (eager semantics); + this rank keeps its chunk of the remapped cells and of the vertex set.""" + mesh = _small_mesh() + window = torch.tensor([2, 3]) # cells [6,7,8], [2,3,9] -> vertices 2,3,6,7,8,9 + ref = _subsample_reference(mesh, window) + for rank in (0, 1): + dm = _StubDeviceMesh(2, rank) + plan = plan_mesh_selection(10, 4, True, True, dm, cell_window=window) + tensors, sharded = mesh_selection_to_proto(mesh, plan) + torch.testing.assert_close(tensors["cells"], ref.cells[rank : rank + 1]) + torch.testing.assert_close( + tensors["points"], ref.points[3 * rank : 3 * rank + 3] + ) + torch.testing.assert_close( + tensors["cell_data", "p"], ref.cell_data["p"][rank : rank + 1] + ) + torch.testing.assert_close( + tensors["cell_data", MEASURE_WEIGHTS_KEY], + ref.cell_data[MEASURE_WEIGHTS_KEY][rank : rank + 1], + ) + assert sharded[("points",)] == (6, 3) and sharded[("cells",)] == (2, 3) + assert plan.measure_factor == 2.0 + with pytest.raises(ValueError, match="point subsampling"): + plan_mesh_selection(10, 4, True, True, dm, point_window=window) + with pytest.raises(ValueError, match="cell_window"): + plan_mesh_selection(10, 0, True, False, dm, cell_window=window) + + +def _subsample_reference(mesh: Mesh, window: torch.Tensor) -> Mesh: + """Eager cell subsample: slice_cells + compaction + measure weights.""" + from physicsnemo.mesh.calculus.measure import compose_measure_weights + + out = mesh.slice_cells(window) + referenced = torch.unique(out.cells) + out = out.slice_points(referenced) + compose_measure_weights(out, mesh.n_cells / len(window)) + return out + + +def test_assemble_if_proto_and_rebuild_registry(): + """Plain samples pass through; an unknown kind fails loudly.""" + plain = object() + assert assemble_if_proto(plain) is plain + proto = ShardedProto( + tensors=TensorDict({"x": torch.zeros(2)}, batch_size=[]), + sharded={}, + device_mesh=None, + kind="nonexistent", + ) + with pytest.raises(ValueError, match="nonexistent"): + proto.assemble() + # With nothing sharded the wrap is a pure copy; the tensordict kind is + # the identity rebuild. + out = dataclasses.replace(proto, kind="tensordict").assemble() + torch.testing.assert_close(out["x"], torch.zeros(2)) + + +def test_leaf_placements_pin_by_prefix_and_warn_unmatched(caplog): + """A dotted override pins every leaf beneath it; a key matching nothing warns.""" + mesh = _StubMesh(2) + meta = { + ("solution", "pressure"): (5000, 1), + ("solution", "velocity"): (5000, 3), + "coords": (5000, 3), + "params": (7,), + } + decisions = resolve_leaf_placements( + meta, + DomainParallelConfig.from_dict({"placements": {"solution": "replicate"}}, mesh), + ) + # The pin reaches both nested leaves, and through the shared axis, coords. + assert decisions[("solution", "pressure")] is False + assert decisions[("solution", "velocity")] is False + assert decisions["coords"] is False + + with caplog.at_level(logging.WARNING): + resolve_leaf_placements( + meta, + DomainParallelConfig.from_dict({"placements": {"coordz": "shard"}}, mesh), + ) + assert any("coordz" in rec.getMessage() for rec in caplog.records) + + +class _NoDomainParallelReader(Reader): + """A reader that does not implement rank-local reading.""" + + def _load_sample(self, index): + return {} + + def __len__(self): + return 1 + + +def test_unsupported_reader_rejects_domain_parallel(): + """Asking a non-supporting reader for domain-parallel reading is an error, + not a silent full read on every rank.""" + with pytest.raises(ValueError, match="does not support domain-parallel"): + _NoDomainParallelReader(domain_parallel={}, device_mesh=_StubMesh(2)) + _NoDomainParallelReader() # no domain parallelism: fine diff --git a/test/domain_parallel/datapipes/__init__.py b/test/domain_parallel/datapipes/__init__.py new file mode 100644 index 0000000000..af85283aa4 --- /dev/null +++ b/test/domain_parallel/datapipes/__init__.py @@ -0,0 +1,15 @@ +# 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. diff --git a/test/domain_parallel/datapipes/test_sharded_domain_mesh_dataset.py b/test/domain_parallel/datapipes/test_sharded_domain_mesh_dataset.py new file mode 100644 index 0000000000..173caf1979 --- /dev/null +++ b/test/domain_parallel/datapipes/test_sharded_domain_mesh_dataset.py @@ -0,0 +1,392 @@ +# 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"""Rank-local sharded reading of DomainMesh through ``MeshDataset``. + +Same skeleton as ``test_sharded_mesh_dataset.py``, over ``.pdmsh`` samples: +interior and a large boundary shard, a tiny boundary is size-gated to a +plain replicated mesh, and the ``apply_to_domain`` transform path is +value-checked against the unsharded reference. +""" + +import pytest +import torch +import torch.distributed as dist +from torch.distributed.tensor.placement_types import Shard + +from physicsnemo.datapipes import MeshDataset +from physicsnemo.datapipes.readers.mesh import DomainMeshReader +from physicsnemo.datapipes.transforms.mesh.transforms import ( + CenterMesh, + NormalizeMeshFields, +) +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel import ShardTensor +from physicsnemo.mesh import DomainMesh, Mesh + +pytestmark = [pytest.mark.multigpu_static, pytest.mark.timeout(300)] + +# Uneven on 2/4/8 ranks for both batch dims of the sharded sub-meshes. +_N_INTERIOR_CELLS = 137 +_N_SURFACE_CELLS = 61 + + +def _triangle_soup(n_cells: int, seed: int, x_offset: float = 0.0) -> Mesh: + r"""One well-conditioned triangle per cell, copies offset along x.""" + torch.manual_seed(seed) + base = torch.tensor([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]) + offsets = torch.zeros(n_cells, 1, 3) + offsets[:, 0, 0] = x_offset + 2.0 * torch.arange(n_cells) + points = (base.unsqueeze(0) + offsets).reshape(-1, 3) + points = points + 0.01 * torch.randn_like(points) + cells = torch.arange(3 * n_cells, dtype=torch.int64).reshape(-1, 3) + return Mesh(points=points, cells=cells) + + +def _build_full_domain(sample: int) -> DomainMesh: + r"""Seeded domain, identical on all ranks: sizeable interior + one large + boundary (shards) + one 3-point probe boundary (size-gated, replicates).""" + interior = _triangle_soup(_N_INTERIOR_CELLS, seed=211 + sample) + interior.point_data["temperature"] = torch.randn(interior.n_points) + interior.cell_data["pressure"] = torch.randn(interior.n_cells) + + surface = _triangle_soup(_N_SURFACE_CELLS, seed=307 + sample, x_offset=500.0) + surface.cell_data["wss"] = torch.randn(surface.n_cells, 3) + + probe = _triangle_soup(1, seed=401 + sample, x_offset=-500.0) + probe.cell_data["flux"] = torch.randn(probe.n_cells) + + return DomainMesh( + interior=interior, + boundaries={"surface": surface, "probe": probe}, + # Nested global data: the domain-parallel read must keep the leaves. + global_data={ + "Re": torch.tensor(1.0e6), + "AoA": torch.tensor(5.0), + "inlet": {"U": torch.tensor([10.0, 0.0, 0.0])}, + }, + ) + + +def _build_extra_boundary(sample: int) -> Mesh: + r"""A sibling surface (the STL of the recipe), sized to pass any gate.""" + return _triangle_soup(_N_SURFACE_CELLS, seed=503 + sample, x_offset=900.0) + + +@pytest.fixture(scope="module") +def pdmsh_root(tmp_path_factory, distributed_mesh): + r"""Shared directory of .pdmsh samples; rank 0 writes, path broadcast.""" + dm = DistributedManager() + if dm.rank == 0: + root = tmp_path_factory.mktemp("sharded_pdmsh") + for i in range(2): + case = root / f"case_{i}" + case.mkdir() + _build_full_domain(i).save(case / f"sample_{i}.pdmsh") + _build_extra_boundary(i).save(case / f"sample_{i}_stl.pmsh") + holder = [str(root)] + else: + holder = [None] + dist.broadcast_object_list(holder, src=0) + return holder[0] + + +def _make_dataset(pdmsh_root, distributed_mesh, transforms=None, **reader_kwargs): + dm = DistributedManager() + reader_kwargs.setdefault("domain_parallel", {"auto_shard_size": 2}) + return MeshDataset( + DomainMeshReader(pdmsh_root, device_mesh=distributed_mesh, **reader_kwargs), + transforms=transforms, + device=dm.device, + ) + + +def _assert_sharded_submesh(mesh: Mesh, full: Mesh, device) -> None: + """Plain layout: points and cells both Shard(0) chunks, cells global ids.""" + assert isinstance(mesh.points, ShardTensor) + assert isinstance(mesh.cells, ShardTensor) + assert mesh.points._spec.placements == (Shard(0),) + assert mesh.n_points == full.n_points + assert mesh.n_cells == full.n_cells + torch.testing.assert_close(mesh.points.full_tensor(), full.points.to(device)) + torch.testing.assert_close(mesh.cells.full_tensor(), full.cells.to(device)) + for key, value in full.point_data.items(): + assert isinstance(mesh.point_data[key], ShardTensor) + torch.testing.assert_close(mesh.point_data[key].full_tensor(), value.to(device)) + for key, value in full.cell_data.items(): + assert isinstance(mesh.cell_data[key], ShardTensor) + torch.testing.assert_close(mesh.cell_data[key].full_tensor(), value.to(device)) + torch.testing.assert_close( + mesh.cell_centroids.full_tensor(), full.cell_centroids.to(device) + ) + + +def _assert_replicated_submesh(mesh: Mesh, full: Mesh, device) -> None: + assert not isinstance(mesh.points, ShardTensor) + torch.testing.assert_close(mesh.points, full.points.to(device)) + torch.testing.assert_close(mesh.cells, full.cells.to(device)) + for key, value in full.cell_data.items(): + assert not isinstance(mesh.cell_data[key], ShardTensor) + torch.testing.assert_close(mesh.cell_data[key], value.to(device)) + + +def test_sharded_domain_read_sync_path(pdmsh_root, distributed_mesh): + r"""dataset[i]: interior + large boundary shard, probe boundary is + size-gated to a plain replicated mesh, domain global_data plain.""" + dm = DistributedManager() + dataset = _make_dataset(pdmsh_root, distributed_mesh) + try: + for i in range(2): + domain, metadata = dataset[i] + full = _build_full_domain(i) + + assert isinstance(domain, DomainMesh) + assert domain.boundary_names == ["probe", "surface"] + _assert_sharded_submesh(domain.interior, full.interior, dm.device) + _assert_sharded_submesh( + domain.boundaries["surface"], full.boundaries["surface"], dm.device + ) + _assert_replicated_submesh( + domain.boundaries["probe"], full.boundaries["probe"], dm.device + ) + for key in ("Re", "AoA"): + assert not isinstance(domain.global_data[key], ShardTensor) + torch.testing.assert_close( + domain.global_data[key], full.global_data[key].to(dm.device) + ) + finally: + dataset.close() + + +def test_sharded_domain_read_producer_consumer_path(pdmsh_root, distributed_mesh): + r"""_load_host -> _consume: host payload carries local rows for the + sharded sub-meshes and full rows for the gated one.""" + from physicsnemo.datapipes._domain_parallel import ShardedProto + + dataset = _make_dataset(pdmsh_root, distributed_mesh) + try: + payload = dataset._load_host(0) + assert payload.error is None + assert isinstance(payload.data, ShardedProto) + assert payload.data.kind == "domain_mesh" + sharded = payload.data.sharded + assert sharded[("interior", "points")] == (3 * _N_INTERIOR_CELLS, 3) + assert sharded[("interior", "cells")] == (_N_INTERIOR_CELLS, 3) + assert ("interior", "cell_data", "pressure") in sharded + assert ("boundaries", "surface", "cells") in sharded + assert not any(key[:2] == ("boundaries", "probe") for key in sharded) + interior_cells = payload.data.tensors["interior", "cells"] + assert ( + interior_cells.shape[0] < _N_INTERIOR_CELLS or distributed_mesh.size(0) == 1 + ) + assert interior_cells.device.type == "cpu" + + domain, _ = dataset._consume(payload) + full = _build_full_domain(0) + dm = DistributedManager() + _assert_sharded_submesh(domain.interior, full.interior, dm.device) + _assert_replicated_submesh( + domain.boundaries["probe"], full.boundaries["probe"], dm.device + ) + finally: + dataset.close() + + +def test_sharded_domain_transforms(pdmsh_root, distributed_mesh): + r"""apply_to_domain chain: CenterMesh reduces the COM from the sharded + interior and translates every sub-mesh (including the replicated probe); + NormalizeMeshFields is elementwise on cell_data. Both must match the + unsharded reference.""" + fields = {"pressure": {"type": "scalar", "mean": 101325.0, "std": 250.0}} + + def make_transforms(): + return [ + CenterMesh(use_area_weighting=False), + NormalizeMeshFields(association="cell_data", fields=fields), + ] + + dm = DistributedManager() + dataset = _make_dataset(pdmsh_root, distributed_mesh, transforms=make_transforms()) + try: + domain, _ = dataset[0] + finally: + dataset.close() + + reference = _build_full_domain(0).to(dm.device) + for t in make_transforms(): + if hasattr(t, "to"): + t.to(dm.device) + reference = t.apply_to_domain(reference) + + assert isinstance(domain.interior.points, ShardTensor) + # 1e-4 on points: CenterMesh's COM is a per-rank partial sum resolved + # by an all-reduce; fp32 summation-order jitter vs the single-device + # reference is a few 1e-5 on coordinates spanning O(1e2) units. + torch.testing.assert_close( + domain.interior.points.full_tensor(), + reference.interior.points, + atol=1e-4, + rtol=1e-4, + ) + torch.testing.assert_close( + domain.interior.cell_data["pressure"].full_tensor(), + reference.interior.cell_data["pressure"], + atol=1e-5, + rtol=1e-5, + ) + # The replicated probe must receive the same translation as the + # sharded interior. + torch.testing.assert_close( + domain.boundaries["probe"].points, + reference.boundaries["probe"].points, + atol=1e-4, + rtol=1e-4, + ) + surface = domain.boundaries["surface"].points + torch.testing.assert_close( + surface.full_tensor(), + reference.boundaries["surface"].points, + atol=1e-4, + rtol=1e-4, + ) + + +def test_placement_override(pdmsh_root, distributed_mesh): + r"""``placements`` pins by path. A mesh with cells has one batch axis + (``cells``; its vertices follow), so pinning a sub-mesh shards or + replicates points and cells together. A gate nothing passes plus an + explicit ``interior: shard`` shards exactly the interior.""" + dm = DistributedManager() + dataset = MeshDataset( + DomainMeshReader( + pdmsh_root, + domain_parallel={ + "auto_shard_size": 10**9, + "placements": {"interior": "shard"}, + }, + device_mesh=distributed_mesh, + ), + device=dm.device, + ) + try: + domain, _ = dataset[0] + full = _build_full_domain(0) + _assert_sharded_submesh(domain.interior, full.interior, dm.device) + for name in ("surface", "probe"): + _assert_replicated_submesh( + domain.boundaries[name], full.boundaries[name], dm.device + ) + finally: + dataset.close() + + # The same pin spelled on the cells axis, and a replicate pin beating the gate. + dataset = MeshDataset( + DomainMeshReader( + pdmsh_root, + domain_parallel={ + "auto_shard_size": 2, + "placements": {"interior.cells": "replicate"}, + }, + device_mesh=distributed_mesh, + ), + device=dm.device, + ) + try: + domain, _ = dataset[0] + full = _build_full_domain(0) + _assert_replicated_submesh(domain.interior, full.interior, dm.device) + _assert_sharded_submesh( + domain.boundaries["surface"], full.boundaries["surface"], dm.device + ) + finally: + dataset.close() + + +def test_nested_global_data_is_kept(pdmsh_root, distributed_mesh): + r"""Nested ``global_data`` leaves survive the rank-local read, replicated.""" + dm = DistributedManager() + dataset = _make_dataset(pdmsh_root, distributed_mesh) + try: + domain, _ = dataset[0] + full = _build_full_domain(0) + for key in (("Re",), ("AoA",), ("inlet", "U")): + value = domain.global_data[key] + assert not isinstance(value, ShardTensor) + torch.testing.assert_close(value, full.global_data[key].to(dm.device)) + finally: + dataset.close() + + +def test_extra_boundaries_replicate_unless_overridden(pdmsh_root, distributed_mesh): + r"""Extra boundaries (full-resolution sibling meshes) are pinned to + replicate by the reader; a user ``placements`` entry can shard them.""" + dm = DistributedManager() + extra = {"stl": {"pattern": "*_stl.pmsh"}} + dataset = _make_dataset(pdmsh_root, distributed_mesh, extra_boundaries=extra) + try: + domain, metadata = dataset[0] + assert "stl" in metadata["boundary_names"] + _assert_replicated_submesh( + domain.boundaries["stl"], _build_extra_boundary(0), dm.device + ) + finally: + dataset.close() + + dataset = _make_dataset( + pdmsh_root, + distributed_mesh, + extra_boundaries=extra, + domain_parallel={ + "auto_shard_size": 2, + "placements": {"boundaries.stl": "shard"}, + }, + ) + try: + domain, _ = dataset[0] + _assert_sharded_submesh( + domain.boundaries["stl"], _build_extra_boundary(0), dm.device + ) + finally: + dataset.close() + + +def test_drop_flags_under_domain_parallel(pdmsh_root, distributed_mesh): + r"""``drop_interior_cells`` reads the interior as a sharded point cloud and + ``drop_in_file_boundaries`` skips the stored boundaries entirely.""" + dm = DistributedManager() + dataset = _make_dataset( + pdmsh_root, + distributed_mesh, + drop_interior_cells=True, + drop_in_file_boundaries=True, + ) + try: + domain, metadata = dataset[0] + full = _build_full_domain(0) + assert metadata["boundary_names"] == [] + assert domain.boundary_names == [] + interior = domain.interior + assert interior.n_cells == 0 + assert isinstance(interior.points, ShardTensor) + torch.testing.assert_close( + interior.points.full_tensor(), full.interior.points.to(dm.device) + ) + torch.testing.assert_close( + interior.point_data["temperature"].full_tensor(), + full.interior.point_data["temperature"].to(dm.device), + ) + finally: + dataset.close() diff --git a/test/domain_parallel/datapipes/test_sharded_mesh_dataset.py b/test/domain_parallel/datapipes/test_sharded_mesh_dataset.py new file mode 100644 index 0000000000..ba46190527 --- /dev/null +++ b/test/domain_parallel/datapipes/test_sharded_mesh_dataset.py @@ -0,0 +1,284 @@ +# 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"""Rank-local sharded reading through ``MeshReader(device_mesh=...)``. + +Rank 0 writes seeded ``.pmsh`` samples to a shared tmp dir; every rank then +reads through the datapipe with ``device_mesh`` set and must see a ``Mesh`` +of ``Shard(0)`` ShardTensors whose gathered values match the full on-disk +sample -- construction, both load paths (sync and producer/consumer), and +the recipe-style transform chain against an unsharded reference. +""" + +import pytest +import torch +import torch.distributed as dist +from torch.distributed.tensor.placement_types import Shard + +from physicsnemo.datapipes import MeshDataset +from physicsnemo.datapipes.readers.mesh import MeshReader +from physicsnemo.datapipes.transforms.mesh.transforms import ( + CenterMesh, + NormalizeMeshFields, +) +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel import ShardTensor +from physicsnemo.mesh import Mesh +from physicsnemo.mesh.io import to_zarr + +pytestmark = [pytest.mark.multigpu_static, pytest.mark.timeout(300)] + +# Uneven on 2/4/8 ranks for both batch dims (n_points = 3 * _N_CELLS). +_N_CELLS = 431 +_N_SAMPLES = 2 + + +def _build_full_mesh(sample: int) -> Mesh: + r"""Seeded triangle soup, distinct per sample, identical on all ranks.""" + torch.manual_seed(101 + sample) + base = torch.tensor([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]) + offsets = torch.zeros(_N_CELLS, 1, 3) + offsets[:, 0, 0] = 2.0 * torch.arange(_N_CELLS) + points = (base.unsqueeze(0) + offsets).reshape(-1, 3) + points = points + 0.01 * torch.randn_like(points) + cells = torch.arange(3 * _N_CELLS, dtype=torch.int64).reshape(-1, 3) + mesh = Mesh(points=points, cells=cells) + mesh.point_data["velocity"] = torch.randn(mesh.n_points, 3) + mesh.cell_data["pressure"] = torch.randn(mesh.n_cells) + mesh.cell_data["wss"] = torch.randn(mesh.n_cells, 3) + mesh.global_data["Re"] = torch.tensor(1.0e6) + return mesh + + +@pytest.fixture(scope="module", params=["pmsh", "zarr"]) +def pmsh_root(request, tmp_path_factory, distributed_mesh): + r"""Shared directory of samples in one on-disk format (memmap ``.pmsh`` or + zarr); rank 0 writes, path broadcast. Both formats go through the same + rank-local read plan but different row sources.""" + fmt = request.param + dm = DistributedManager() + if dm.rank == 0: + root = tmp_path_factory.mktemp(f"sharded_{fmt}") + for i in range(_N_SAMPLES): + mesh = _build_full_mesh(i) + if fmt == "pmsh": + mesh.save(root / f"sample_{i}.pmsh") + else: + to_zarr(mesh, root / f"sample_{i}.zarr") + holder = [str(root), fmt] + else: + holder = [None, None] + dist.broadcast_object_list(holder, src=0) + return holder[0], holder[1] + + +def _make_reader(root_and_fmt, distributed_mesh, **kwargs): + root, fmt = root_and_fmt + kwargs.setdefault("domain_parallel", {"auto_shard_size": 2}) + kwargs.setdefault("device_mesh", distributed_mesh) + return MeshReader(root, pattern=f"**/*.{fmt}", **kwargs) + + +def _make_dataset(pmsh_root, distributed_mesh, transforms=None, **reader_kwargs): + dm = DistributedManager() + return MeshDataset( + _make_reader(pmsh_root, distributed_mesh, **reader_kwargs), + transforms=transforms, + device=dm.device, + ) + + +def _assert_sharded_matches_full(mesh, sample: int, distributed_mesh): + """Plain layout: points/point_data chunk over points, cells/cell_data chunk + over cells (cells keep global vertex ids); every gathered leaf equals the + unsharded sample and per-cell quantities go through the routed gather.""" + full = _build_full_mesh(sample) + world_size = distributed_mesh.size(0) + + assert mesh.n_points == full.n_points + assert mesh.n_cells == full.n_cells + for leaf in (mesh.points, mesh.cells): + assert isinstance(leaf, ShardTensor) + assert leaf._spec.placements == (Shard(0),) + assert mesh.points._local_tensor.shape[0] <= -(-full.n_points // world_size) + assert mesh.cells._local_tensor.shape[0] <= -(-full.n_cells // world_size) + + device = mesh.points._local_tensor.device + torch.testing.assert_close(mesh.points.full_tensor(), full.points.to(device)) + torch.testing.assert_close(mesh.cells.full_tensor(), full.cells.to(device)) + torch.testing.assert_close( + mesh.point_data["velocity"].full_tensor(), + full.point_data["velocity"].to(device), + ) + for key in ("pressure", "wss"): + assert isinstance(mesh.cell_data[key], ShardTensor) + torch.testing.assert_close( + mesh.cell_data[key].full_tensor(), full.cell_data[key].to(device) + ) + torch.testing.assert_close( + mesh.global_data["Re"], full.global_data["Re"].to(device) + ) + # Cell quantities: points[cells] is the routed gather across ranks. + centroids = mesh.cell_centroids + assert isinstance(centroids, ShardTensor) + torch.testing.assert_close(centroids.full_tensor(), full.cell_centroids.to(device)) + + +def test_sharded_read_sync_path(pmsh_root, distributed_mesh): + r"""dataset[i] (synchronous _load): sharded Mesh matches the full sample.""" + dataset = _make_dataset(pmsh_root, distributed_mesh) + try: + for i in range(_N_SAMPLES): + mesh, metadata = dataset[i] + _assert_sharded_matches_full(mesh, i, distributed_mesh) + assert metadata["index"] == i + finally: + dataset.close() + + +def test_sharded_read_producer_consumer_path(pmsh_root, distributed_mesh): + r"""_load_host -> _consume (the prefetch stages, no stream): the slice + happens host-side, the ShardTensor wrap after device transfer.""" + dataset = _make_dataset(pmsh_root, distributed_mesh) + try: + payload = dataset._load_host(0) + assert payload.error is None + # Host payload carries this rank's chunks only. + local_cells = payload.data.tensors["cells"].shape[0] + local_points = payload.data.tensors["points"].shape[0] + assert local_cells < _N_CELLS or distributed_mesh.size(0) == 1 + assert local_points < 3 * _N_CELLS or distributed_mesh.size(0) == 1 + assert payload.data.tensors["points"].device.type == "cpu" + assert payload.data.sharded[("points",)] == (3 * _N_CELLS, 3) + assert payload.data.sharded[("cells",)] == (_N_CELLS, 3) + + mesh, _ = dataset._consume(payload) + _assert_sharded_matches_full(mesh, 0, distributed_mesh) + finally: + dataset.close() + + +def test_sharded_read_with_transforms(pmsh_root, distributed_mesh): + r"""Recipe-style transform chain on the sharded pipe matches the same + chain applied to the full mesh: CenterMesh is the global reduction, + NormalizeMeshFields the elementwise cell_data op.""" + fields = { + "pressure": {"type": "scalar", "mean": 101325.0, "std": 250.0}, + "wss": {"type": "vector", "mean": [1.0, 0.0, 0.0], "std": 0.5}, + } + + def make_transforms(): + return [ + CenterMesh(use_area_weighting=False), + NormalizeMeshFields(association="cell_data", fields=fields), + ] + + dm = DistributedManager() + dataset = _make_dataset(pmsh_root, distributed_mesh, transforms=make_transforms()) + try: + mesh, _ = dataset[0] + finally: + dataset.close() + + reference = _build_full_mesh(0).to(dm.device) + for t in make_transforms(): + if hasattr(t, "to"): + t.to(dm.device) + reference = t(reference) + + assert isinstance(mesh.points, ShardTensor) + # 1e-4: CenterMesh's COM is a per-rank partial sum resolved by an + # all-reduce; fp32 summation-order jitter vs the single-device + # reference is a few 1e-5 on coordinates spanning O(1e3) units. + torch.testing.assert_close( + mesh.points.full_tensor(), reference.points, atol=1e-4, rtol=1e-4 + ) + for key in ("pressure", "wss"): + torch.testing.assert_close( + mesh.cell_data[key].full_tensor(), + reference.cell_data[key], + atol=1e-5, + rtol=1e-5, + ) + + +def test_cell_subsample_matches_eager(pmsh_root, distributed_mesh): + r"""``subsample_n_cells`` under domain parallelism: the whole window of + cells is compacted identically on every rank and the gathered sharded + mesh equals the eager (unsharded) subsample with the same seed/epoch.""" + dm = DistributedManager() + root, fmt = pmsh_root + n_cells = 97 + seed, epoch = 4321, 2 + + def seeded(dataset): + generator = torch.Generator() + generator.manual_seed(seed) + dataset.set_generator(generator) + dataset.set_epoch(epoch) + return dataset + + reference_ds = seeded( + MeshDataset( + MeshReader(root, pattern=f"**/*.{fmt}", subsample_n_cells=n_cells), + device=dm.device, + ) + ) + sharded_ds = seeded( + _make_dataset(pmsh_root, distributed_mesh, subsample_n_cells=n_cells) + ) + try: + reference, _ = reference_ds[1] + mesh, _ = sharded_ds[1] + assert mesh.n_cells == n_cells == reference.n_cells + assert mesh.n_points == reference.n_points + torch.testing.assert_close(mesh.cells.full_tensor(), reference.cells) + torch.testing.assert_close(mesh.points.full_tensor(), reference.points) + torch.testing.assert_close( + mesh.cell_data["wss"].full_tensor(), reference.cell_data["wss"] + ) + torch.testing.assert_close( + mesh.point_data["velocity"].full_tensor(), + reference.point_data["velocity"], + ) + finally: + reference_ds.close() + sharded_ds.close() + + +def test_point_subsample_on_cells_is_rejected(pmsh_root, distributed_mesh): + r"""Point subsampling on a mesh with cells is not supported under domain + parallelism (it would remap connectivity globally); the reader raises + before any collective, identically on every rank.""" + dataset = _make_dataset(pmsh_root, distributed_mesh, subsample_n_points=50) + generator = torch.Generator() + generator.manual_seed(0) + dataset.set_generator(generator) + try: + with pytest.raises(NotImplementedError, match="subsample_n_points"): + dataset[0] + finally: + dataset.close() + + +def test_subsample_without_seed_is_rejected(pmsh_root, distributed_mesh): + r"""An unseeded window would differ per rank; the reader refuses up front.""" + dataset = _make_dataset(pmsh_root, distributed_mesh, subsample_n_cells=97) + try: + with pytest.raises(ValueError, match="requires a seed"): + dataset[0] + finally: + dataset.close() diff --git a/test/domain_parallel/datapipes/test_sharded_tensordict_dataset.py b/test/domain_parallel/datapipes/test_sharded_tensordict_dataset.py new file mode 100644 index 0000000000..8157d54c02 --- /dev/null +++ b/test/domain_parallel/datapipes/test_sharded_tensordict_dataset.py @@ -0,0 +1,204 @@ +# 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"""Generic domain-parallel reading: ZarrReader -> Dataset -> ShardTensors. + +Rank 0 writes seeded zarr groups to a shared dir; every rank reads through +``ZarrReader(domain_parallel=..., device_mesh=...)`` + ``Dataset`` and must +see per-key ``Shard(0)`` ShardTensors (or plain replicated tensors) whose +gathered values match the on-disk sample -- in manual and auto placement +modes, and composed with coordinated subsampling. +""" + +import numpy as np +import pytest +import torch +import torch.distributed as dist +from torch.distributed.tensor.placement_types import Shard + +from physicsnemo.datapipes import Dataset +from physicsnemo.datapipes.readers.zarr import ZarrReader +from physicsnemo.distributed import DistributedManager +from physicsnemo.domain_parallel import ShardTensor + +zarr = pytest.importorskip("zarr") + +pytestmark = [pytest.mark.multigpu_static, pytest.mark.timeout(300)] + +# Uneven on 2/4/8 ranks. +_N_ROWS = 1234 +_N_SAMPLES = 2 + + +def _sample_arrays(sample: int) -> dict[str, np.ndarray]: + rng = np.random.default_rng(97 + sample) + return { + "coords": rng.standard_normal((_N_ROWS, 3), dtype=np.float32), + "fields": rng.standard_normal((_N_ROWS, 4), dtype=np.float32), + "params": rng.standard_normal((7,), dtype=np.float32), + } + + +@pytest.fixture(scope="module") +def zarr_root(tmp_path_factory, distributed_mesh): + r"""Shared directory of zarr groups; rank 0 writes, path broadcast.""" + dm = DistributedManager() + if dm.rank == 0: + root = tmp_path_factory.mktemp("sharded_zarr") + for i in range(_N_SAMPLES): + group = zarr.open_group(str(root / f"sample_{i}.zarr"), mode="w") + for key, value in _sample_arrays(i).items(): + group[key] = value + holder = [str(root)] + else: + holder = [None] + dist.broadcast_object_list(holder, src=0) + return holder[0] + + +def _make_dataset(zarr_root, distributed_mesh, **reader_kwargs): + dm = DistributedManager() + return Dataset( + ZarrReader(zarr_root, device_mesh=distributed_mesh, **reader_kwargs), + device=dm.device, + ) + + +def _assert_sharded_key(td, key, reference, device): + assert isinstance(td[key], ShardTensor) + assert td[key]._spec.placements == (Shard(0),) + torch.testing.assert_close( + td[key].full_tensor(), torch.from_numpy(reference).to(device) + ) + + +def _assert_replicated_key(td, key, reference, device): + assert not isinstance(td[key], ShardTensor) + torch.testing.assert_close(td[key], torch.from_numpy(reference).to(device)) + + +@pytest.mark.parametrize("pin_memory", [False, True]) +def test_manual_placements(zarr_root, distributed_mesh, pin_memory): + r"""A pinned key shards its whole batch axis (``fields`` shares the + 1234-row axis with ``coords``); unpinned axes fall back to the gate, which + the huge ``auto_shard_size`` fails, so ``params`` replicates.""" + dm = DistributedManager() + dataset = _make_dataset( + zarr_root, + distributed_mesh, + pin_memory=pin_memory, + domain_parallel={ + "auto_shard_size": 10**9, + "placements": {"coords": "shard"}, + }, + ) + try: + for i in range(_N_SAMPLES): + td, metadata = dataset[i] + reference = _sample_arrays(i) + _assert_sharded_key(td, "coords", reference["coords"], dm.device) + _assert_sharded_key(td, "fields", reference["fields"], dm.device) + _assert_replicated_key(td, "params", reference["params"], dm.device) + finally: + dataset.close() + + +def test_gate_decision_reaches_the_wrap(zarr_root, distributed_mesh): + r"""The auto_shard_size gate decides per axis and the decision is what + the dataset materializes: coords/fields (1234 long) shard at a gate of + 100 while params (7 long) replicates; a gate nothing passes replicates + all. The gate arithmetic itself is covered single-rank in + ``test/datapipes/test_domain_parallel_config.py``.""" + dm = DistributedManager() + dataset = _make_dataset( + zarr_root, distributed_mesh, domain_parallel={"auto_shard_size": 100} + ) + try: + td, _ = dataset[0] + reference = _sample_arrays(0) + _assert_sharded_key(td, "coords", reference["coords"], dm.device) + _assert_sharded_key(td, "fields", reference["fields"], dm.device) + _assert_replicated_key(td, "params", reference["params"], dm.device) + finally: + dataset.close() + + dataset = _make_dataset( + zarr_root, distributed_mesh, domain_parallel={"auto_shard_size": 10**9} + ) + try: + td, _ = dataset[0] + reference = _sample_arrays(0) + for key in ("coords", "fields", "params"): + _assert_replicated_key(td, key, reference[key], dm.device) + finally: + dataset.close() + + +def test_sharded_read_composes_with_subsampling(zarr_root, distributed_mesh): + r"""The rank chunk is taken OF the coordinated window: the gathered + sharded read equals a full (unsharded) read with the same seed/epoch.""" + dm = DistributedManager() + n_window = 600 + subsampling = {"n_points": n_window, "target_keys": ["coords", "fields"]} + seed = 1234 + + def build(domain_parallel=None): + reader_kwargs = {"coordinated_subsampling": subsampling} + if domain_parallel is not None: + reader_kwargs["domain_parallel"] = domain_parallel + reader_kwargs["device_mesh"] = distributed_mesh + dataset = Dataset(ZarrReader(zarr_root, **reader_kwargs), device=dm.device) + generator = torch.Generator() + generator.manual_seed(seed) + dataset.set_generator(generator) + dataset.set_epoch(3) + return dataset + + reference_ds = build() + sharded_ds = build(domain_parallel={"auto_shard_size": 100}) + try: + reference, _ = reference_ds[1] + td, _ = sharded_ds[1] + + for key in ("coords", "fields"): + assert isinstance(td[key], ShardTensor) + assert td[key].shape[0] == n_window + torch.testing.assert_close(td[key].full_tensor(), reference[key]) + # Replicated keys are identical to the reference read. + torch.testing.assert_close(td["params"], reference["params"]) + finally: + reference_ds.close() + sharded_ds.close() + + +def test_subsampling_without_seed_raises(zarr_root, distributed_mesh): + r"""Coordinated subsampling under domain parallelism needs a seed, or the + ranks would draw different windows; the reader refuses up front.""" + dm = DistributedManager() + dataset = Dataset( + ZarrReader( + zarr_root, + coordinated_subsampling={"n_points": 600, "target_keys": ["coords"]}, + domain_parallel={"auto_shard_size": 100}, + device_mesh=distributed_mesh, + ), + device=dm.device, + ) + try: + with pytest.raises(ValueError, match="requires a seed"): + dataset[0] + finally: + dataset.close()