From 7a914e1c70b52a54d7fc8e10dbafc6ac03b86412 Mon Sep 17 00:00:00 2001 From: Corey Adams <6619961+coreyjadams@users.noreply.github.com> Date: Thu, 10 Sep 2026 16:22:45 +0000 Subject: [PATCH 1/2] Domain parallelism in the unified external aero recipe `domain_parallelism.domain_size=N` shards every sample over N GPUs; the remaining ranks form the data-parallel axis. Samples arrive from the readers already sharded (level 4) and the model runs on ShardTensors, with DDP over the flat world for the parameter gradients. `domain_size=1` leaves every code path identical to a non-domain-parallel run. - conf/base.yaml: `domain_parallelism` block (`domain_size`, `auto_shard_size`, `placements`); everything but `domain_size` is the readers' `domain_parallel` policy and is passed through untouched. - utils.build_distributed_meshes: builds the ("ddp", "domain") device mesh, checks divisibility, returns (domain_mesh, data_mesh) or (None, None). - datasets.build_dataloaders: keyword-only `domain_mesh` / `data_mesh`; readers get `device_mesh` + the policy, directory and manifest samplers shard over the data-parallel axis only so every rank of a domain group sees the same sample sequence. - train.py: RNG seeded by data-parallel rank instead of world rank; `sync_module_over_mesh` before DDP so the domain group starts from one set of parameters; `materialize` resolves sharded 0-D losses/metrics (recursing into TensorDicts) before logging and the benchmark stats. - README: short "Domain parallelism" section pointing at the guide. - Tests: CPU tests for the config pass-through and the sampler sharding over a fake data mesh; GPU test comparing loss and parameter gradients of a sync + DDP domain-parallel step against a single process. Known gaps, fixed on the lower stack levels: FLARE under bf16 autocast gives NaN gradients through ReplicatedQSDPA, and `compile=True` with `domain_size > 1` fails in the sharded radius search. --- CHANGELOG.md | 3 + .../unified_external_aero_recipe/README.md | 24 +++ .../conf/base.yaml | 8 + .../src/datasets.py | 89 +++++++++- .../unified_external_aero_recipe/src/train.py | 55 +++++- .../unified_external_aero_recipe/src/utils.py | 31 +++- .../tests/test_domain_parallel.py | 130 ++++++++++++++ .../tests/test_domain_parallel_gradients.py | 159 ++++++++++++++++++ 8 files changed, 483 insertions(+), 16 deletions(-) create mode 100644 examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel.py create mode 100644 examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel_gradients.py diff --git a/CHANGELOG.md b/CHANGELOG.md index f16c710f75..406dbc4e3e 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,6 +32,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 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`. +- The unified external aero recipe trains with optoinal domain parallelism + (`domain_parallelism.domain_size`), reading samples sharded and running the + model on ShardTensors over a domain x data-parallel device mesh. ### Changed diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/README.md b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/README.md index c0bc75a9fd..d71c04857a 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/README.md +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/README.md @@ -405,6 +405,30 @@ python src/train.py benchmark_io=true +training.benchmark_max_steps=20 Measures per-sample load time and throughput without running the model. +### Domain parallelism + +The unified external aero recipe supports domain parallelism for select models. +This recipe uses ``ShardTensor`` in physicsnemo, supported natively in physicsnemo's +datapipes and mesh objects, to perform domain parallel IO + preprocessing as well +as model training. + +Set `domain_parallelism.domain_size=N` to shard each sample across N GPUs; +the remaining GPUs form the data-parallel (DDP) axis (so `world_size` must be +a multiple of N). + +```bash +torchrun --nproc-per-node 4 src/train.py model=geotransolver_volume \ + dataset=drivaer_ml_volume domain_parallelism.domain_size=4 +``` + +The other keys in the `domain_parallelism` block (`auto_shard_size`, +`placements`) are the readers' sharding policy; see the +[Domain Parallelism Guide](https://docs.nvidia.com/deeplearning/physicsnemo/physicsnemo-core/tutorials/domain_parallelism_entry_point.html) +for more information on domain parallelism. + +In the recipe, both surface and volume training are supported for domain parallelism. +Transolver, GeoTransolver, and FLARE are supported. GLOBE is not yet supported. + ## Configuration A single canonical `conf/train.yaml` drives every training run. It picks diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/conf/base.yaml b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/conf/base.yaml index 052bf8f8a8..d01915d25b 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/conf/base.yaml +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/conf/base.yaml @@ -49,6 +49,14 @@ dataloader: num_workers: 4 pin_memory: true +# -- Domain parallelism ------------------------------------------------------- +# Shard each sample over `domain_size` GPUs; the rest of the block is the +# readers' `domain_parallel` policy (see physicsnemo.datapipes.readers.Reader). +domain_parallelism: + domain_size: 1 # number of gpus per batch + auto_shard_size: 10000 + placements: {} + # -- Logging ----------------------------------------------------------------- logging: log_every_n_steps: 10 diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/datasets.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/datasets.py index 55701642a8..63195687e6 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/datasets.py +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/datasets.py @@ -41,7 +41,7 @@ import sys from collections.abc import Callable, Iterator from pathlib import Path -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast import hydra import torch @@ -53,6 +53,9 @@ from physicsnemo.datapipes.transforms.mesh import NormalizeMeshFields from physicsnemo.distributed import DistributedManager +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh + ### Make this folder importable by its bare module names (`nondim`, `sdf`) ### regardless of whether the caller invoked `python src/train.py` (which ### already adds `src/` to sys.path[0]) or imported `datasets` from a @@ -330,6 +333,8 @@ def build_dataset( device: str | torch.device | None = "auto", num_workers: int = 1, pin_memory: bool = False, + domain_parallel: dict | None = None, + device_mesh: "DeviceMesh | None" = None, ) -> MeshDataset: """Build a single MeshDataset from a Hydra-style pipeline config. @@ -358,6 +363,12 @@ def build_dataset( prefetch pool. pin_memory: If True, the reader places tensors in pinned (page-locked) memory for faster async CPU-to-GPU transfers. + domain_parallel: Optional domain-parallel reading policy dict + (see :class:`physicsnemo.datapipes.readers.base.Reader`); + requires *device_mesh*. The reader then reads each rank's + share and the dataset assembles ShardTensor-backed meshes. + device_mesh: Runtime 1-D DeviceMesh for domain-parallel reading + (the ``domain`` axis of the training mesh); never config. Returns: Configured ``MeshDataset`` ready to be wrapped in a DataLoader. @@ -365,7 +376,12 @@ def build_dataset( if base_dir is None: base_dir = Path(__file__).resolve().parent.parent - reader = hydra.utils.instantiate(cfg.pipeline.reader, pin_memory=pin_memory) + ### The domain-parallel policy is already a plain dict (see + ### ``_domain_parallel_from_cfg``), so both branches instantiate the same way. + reader_kwargs: dict[str, Any] = {"pin_memory": pin_memory} + if domain_parallel is not None: + reader_kwargs.update(domain_parallel=domain_parallel, device_mesh=device_mesh) + reader = hydra.utils.instantiate(cfg.pipeline.reader, **reader_kwargs) resolved = [] target_names = list( @@ -622,6 +638,8 @@ def _build_manifest_val_dataset( device: str | torch.device | None, num_workers: int, pin_memory: bool, + domain_parallel: dict | None = None, + device_mesh: "DeviceMesh | None" = None, ) -> MeshDataset | None: """Build a dedicated un-augmented validation dataset for manifest mode. @@ -647,6 +665,8 @@ def _build_manifest_val_dataset( device=device, num_workers=num_workers, pin_memory=pin_memory, + domain_parallel=domain_parallel, + device_mesh=device_mesh, ) @@ -690,6 +710,7 @@ def _build_directory_samplers( *, use_distributed: bool, sampler_seed: int, + data_mesh: "DeviceMesh | None" = None, ) -> tuple[Sampler | None, Sampler | None]: """Per-split :class:`DistributedSampler` pair for **directory-mode** datasets. @@ -698,14 +719,24 @@ def _build_directory_samplers( single dataset across splits and uses :func:`_build_manifest_samplers` instead. Returns ``(None, None)`` on a single rank, where torch's default sequential sampler is sufficient. + + Under domain parallelism, samples are distributed over the ``ddp`` + axis only (*data_mesh*): every rank in a domain group must receive + the same sample index, since each holds one shard of that sample. """ if not use_distributed: return None, None + replica_kwargs = {} + if data_mesh is not None: + replica_kwargs = { + "num_replicas": data_mesh.size(), + "rank": data_mesh.get_local_rank(), + } train_sampler = torch.utils.data.distributed.DistributedSampler( - train_dataset, shuffle=True, drop_last=True, seed=sampler_seed + train_dataset, shuffle=True, drop_last=True, seed=sampler_seed, **replica_kwargs ) val_sampler = torch.utils.data.distributed.DistributedSampler( - val_dataset, shuffle=False, drop_last=False + val_dataset, shuffle=False, drop_last=False, **replica_kwargs ) return train_sampler, val_sampler @@ -716,11 +747,21 @@ def _build_manifest_samplers( *, dist_manager: DistributedManager, sampler_seed: int, + data_mesh: "DeviceMesh | None" = None, ) -> tuple[ManifestSampler, ManifestSampler]: - """ManifestSamplers over global dataset indices, with optional sharding.""" - use_distributed = dist_manager.world_size > 1 - rank = dist_manager.rank if use_distributed else 0 - world_size = dist_manager.world_size if use_distributed else 1 + """ManifestSamplers over global dataset indices, with optional sharding. + + Under domain parallelism, indices are sharded over the ``ddp`` axis + only (*data_mesh*): every rank in a domain group must receive the same + sample index, since each holds one shard of that sample. + """ + if data_mesh is not None: + rank = data_mesh.get_local_rank() + world_size = data_mesh.size() + elif dist_manager.world_size > 1: + rank, world_size = dist_manager.rank, dist_manager.world_size + else: + rank, world_size = 0, 1 train_sampler = ManifestSampler( train_indices, @@ -741,8 +782,26 @@ def _build_manifest_samplers( return train_sampler, val_sampler +def _domain_parallel_from_cfg(cfg: DictConfig) -> dict[str, Any]: + """The reader-facing ``domain_parallel`` dict from ``cfg.domain_parallelism``. + + Everything in the block except ``domain_size`` (a launch-topology knob + consumed by ``train.py``) is handed to the readers unchanged, so the + datapipes' own validation sees exactly what the user wrote. + """ + block = OmegaConf.to_container( + cfg.get("domain_parallelism", OmegaConf.create({})), resolve=True + ) + policy = dict(block or {}) + policy.pop("domain_size", None) + return policy + + def build_dataloaders( cfg: DictConfig, + *, + domain_mesh: "DeviceMesh | None" = None, + data_mesh: "DeviceMesh | None" = None, ) -> tuple[DataLoader, DataLoader, "NormalizeMeshFields | None", dict[str, Any]]: """Build train and val dataloaders from the chosen dataset(s). @@ -800,6 +859,12 @@ def build_dataloaders( device = "cuda" if torch.cuda.is_available() else "cpu" sampler_seed = cfg.training.get("seed", 0) or 0 + ### Domain parallelism: declarative reading policy from cfg, paired + ### with the runtime device mesh injected by the caller. None when off. + domain_parallel = ( + _domain_parallel_from_cfg(cfg) if domain_mesh is not None else None + ) + ### The primary dataset is `cfg.dataset` (a single string); extras ### combine via MultiDataset. The same `train_split`/`val_split` ### apply to every chosen dataset; when they are set, @@ -897,6 +962,8 @@ def build_dataloaders( device=device, num_workers=num_workers, pin_memory=pin_memory, + domain_parallel=domain_parallel, + device_mesh=domain_mesh, ) train_datasets.append(dataset) @@ -936,6 +1003,8 @@ def build_dataloaders( device=device, num_workers=num_workers, pin_memory=pin_memory, + domain_parallel=domain_parallel, + device_mesh=domain_mesh, ) val_dataset = ( manifest_val_dataset if manifest_val_dataset is not None else dataset @@ -959,6 +1028,8 @@ def build_dataloaders( device=device, num_workers=num_workers, pin_memory=pin_memory, + domain_parallel=domain_parallel, + device_mesh=domain_mesh, ) val_datasets.append(val_dataset) combined_val_indices.extend( @@ -997,6 +1068,7 @@ def build_dataloaders( combined_val_indices, dist_manager=dist_manager, sampler_seed=sampler_seed, + data_mesh=data_mesh, ) else: ### Directory mode: separate datasets per split, with per-rank @@ -1007,6 +1079,7 @@ def build_dataloaders( val_dataset, use_distributed=use_distributed, sampler_seed=sampler_seed, + data_mesh=data_mesh, ) ### Shared loader knobs; the two splits differ only in dataset / shuffle / diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py index 9964eb5567..99d2c2c02f 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py @@ -56,6 +56,7 @@ FieldType, Phase, Precision, + build_distributed_meshes, build_muon_optimizer, get_autocast_context, make_jsonl_logger, @@ -67,6 +68,7 @@ from physicsnemo import datapipes # noqa: F401 - registers ${dp:...} resolver from physicsnemo.datapipes import DataLoader from physicsnemo.distributed import DistributedManager, fused_all_reduce +from physicsnemo.domain_parallel import sync_module_over_mesh from physicsnemo.mesh import MESH_FIELD_ASSOCIATIONS, DomainMesh, Mesh from physicsnemo.utils import load_checkpoint, save_checkpoint from physicsnemo.utils.logging import PythonLogger, RankZeroLoggingWrapper @@ -102,6 +104,17 @@ def _flatten_config( ### --------------------------------------------------------------------------- +def materialize(t: torch.Tensor | TensorDict) -> torch.Tensor | TensorDict: + """Resolve sharded 0-D results to plain tensors, recursing into TensorDicts. + + ShardTensor leaves are gathered with ``full_tensor()``; plain tensors pass + through. + """ + if isinstance(t, TensorDict): + return TensorDict({key: materialize(value) for key, value in t.items()}) + return t.full_tensor() if hasattr(t, "full_tensor") else t + + def _reduce_and_average( loss_sum: Float[torch.Tensor, ""], losses_td: TensorDict | None, @@ -160,6 +173,10 @@ def _reduce_and_average( same leaves in the same order (all ranks share one ``target_config``). Single-process skips the reduction, leaving single-GPU logs unchanged. """ + ### Inputs are plain tensors even under domain parallelism: ``forward_pass`` + ### materializes the logging TensorDicts and ``_run_epoch`` the detached + ### loss, so values are identical across each domain group and the + ### world-wide AVG below equals the ddp-axis mean. if losses_td is None or metrics_td is None: return loss_sum.item() / max(n_samples, 1), {}, {} ### Divide by the local sample count first, then AVG across ranks: a @@ -268,8 +285,15 @@ def forward_pass( ### Detach (don't sync) the per-field TDs so the caller controls when ### a D2H copy happens; running ``.item()`` here would serialise the ### forward kernels against the host. ``TensorDict.detach()`` walks - ### every leaf in one fast-apply pass. - return loss, loss_td.detach(), metric_td.detach() + ### every leaf in one fast-apply pass. Under domain parallelism the + ### leaves are ShardTensors; the detached LOGGING copies are materialized + ### to plain tensors while the live ``loss`` stays untouched for + ### ``backward()``. + return ( + loss, + materialize(loss_td.detach()), + materialize(metric_td.detach()), + ) ### --------------------------------------------------------------------------- @@ -384,8 +408,9 @@ def _run_epoch( n_local += 1 ### Detached scalar loss: accumulate the epoch sum on-device (no - ### host sync) and feed the per-step reducer below. - loss_det = loss.detach() + ### host sync) and feed the per-step reducer below. Materialized + ### so the plain accumulator never sees a ShardTensor. + loss_det = materialize(loss.detach()) total_loss += loss_det step_dt = time.perf_counter() - step_t0 @@ -692,6 +717,7 @@ def benchmark_io_epoch( f"dt={dt:.4f}s Mem={mem_gb:.2f}GB {shapes}" ) for name, t in named_tensors: + t = materialize(t) # sharded leaves: stats over the full tensor v_flat = t.float() if t.is_floating_point() else t.to(torch.float32) logger.info( f" {name:30s} " @@ -753,9 +779,16 @@ def main(cfg: DictConfig) -> None: is_rank0 = dist_manager.rank == 0 logger = RankZeroLoggingWrapper(PythonLogger(name="training"), dist_manager) + domain_mesh, data_mesh = build_distributed_meshes(cfg, dist_manager, logger) + + ### Seed per data-parallel replica: every rank of a domain group works on + ### the same sample and must draw identical dropout / augmentation RNG. seed = cfg.training.get("seed", None) - set_seed(seed, rank=dist_manager.rank) - logger.info(f"Random seed: {seed} (rank offset: {dist_manager.rank})") + seed_rank = ( + data_mesh.get_local_rank() if data_mesh is not None else dist_manager.rank + ) + set_seed(seed, rank=seed_rank) + logger.info(f"Random seed: {seed} (rank offset: {seed_rank})") checkpoint_dir = getattr(cfg, "checkpoint_dir", None) or cfg.output_dir @@ -772,7 +805,9 @@ def main(cfg: DictConfig) -> None: val_writer = SummaryWriter(log_dir=os.path.join(run_dir, "tb", "val")) log_jsonl = make_jsonl_logger(os.path.join(run_dir, "metrics.jsonl")) - train_loader, val_loader, normalizer, dataset_info = build_dataloaders(cfg) + train_loader, val_loader, normalizer, dataset_info = build_dataloaders( + cfg, domain_mesh=domain_mesh, data_mesh=data_mesh + ) target_config: dict[str, FieldType] = dataset_info["targets"] ### `metrics_list` is derived later from cfg.metrics (recipe-side); ### build_dataloaders no longer ships a "metrics" key in dataset_info. @@ -834,6 +869,12 @@ def main(cfg: DictConfig) -> None: model.to(device) + if domain_mesh is not None: + # All ranks in a domain group process the same sample and must hold + # identical weights; DDP below all-reduces gradients over the flat + # world, which covers both the ddp and domain axes. + sync_module_over_mesh(model, domain_mesh) + if dist_manager.world_size > 1: model = torch.nn.parallel.DistributedDataParallel( model, diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/utils.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/utils.py index 2436c1a4eb..9c31146dae 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/utils.py +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/utils.py @@ -24,7 +24,7 @@ from contextlib import nullcontext from datetime import datetime, timezone from pathlib import Path -from typing import Any, Literal, TypeAlias +from typing import TYPE_CHECKING, Any, Literal, TypeAlias import numpy as np import torch @@ -33,9 +33,13 @@ from torch.amp import autocast from physicsnemo.datapipes.keys import as_nested_key +from physicsnemo.distributed import DistributedManager from physicsnemo.mesh import DomainMesh, Mesh from physicsnemo.optim import CombinedOptimizer, Muon +if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh + ### Recipe-wide type aliases. Re-exported for use in loss.py, metrics.py, ### output_normalize.py, forward_kwargs.py, collate.py, train.py, infer.py, ### and the tests so that ``target_config`` values share a single source of @@ -83,6 +87,31 @@ def set_seed(seed: int | None, rank: int = 0) -> None: torch.cuda.manual_seed_all(seed) +def build_distributed_meshes( + cfg: DictConfig, dist_manager: DistributedManager, logger: Any +) -> tuple["DeviceMesh | None", "DeviceMesh | None"]: + """Build the ``(domain_mesh, data_mesh)`` pair for domain parallelism. + + Both are ``None`` when ``domain_size == 1``. The ``domain`` axis varies + fastest, so a domain group is a block of consecutive ranks. + """ + domain_size = int(cfg.get("domain_parallelism", {}).get("domain_size", 1)) + if domain_size <= 1: + return None, None + if dist_manager.world_size % domain_size != 0: + raise ValueError( + f"world_size {dist_manager.world_size} is not divisible by " + f"domain_parallelism.domain_size {domain_size}" + ) + global_mesh = dist_manager.initialize_mesh( + mesh_shape=(-1, domain_size), mesh_dim_names=("ddp", "domain") + ) + domain_mesh = global_mesh["domain"] + data_mesh = global_mesh["ddp"] + logger.info(f"Domain parallelism: ddp={data_mesh.size()} x domain={domain_size}") + return domain_mesh, data_mesh + + def build_muon_optimizer( model: torch.nn.Module, cfg: DictConfig, *, compile_optimizer: bool = False ) -> torch.optim.Optimizer: diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel.py new file mode 100644 index 0000000000..b2b5a40dc5 --- /dev/null +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel.py @@ -0,0 +1,130 @@ +# 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. + +"""CPU tests for the recipe's domain-parallel plumbing. + +Config translation to the readers and sampling over the data-parallel axis +only. The gradient-equivalence check across a domain group lives in +``test_domain_parallel_gradients.py`` (multi-GPU). +""" + +from __future__ import annotations + +from datasets import ( + _build_directory_samplers, + _build_manifest_samplers, + _domain_parallel_from_cfg, +) +from omegaconf import OmegaConf + + +class _FakeMesh: + """Stand-in for a 1-D DeviceMesh axis: a size and this rank's position.""" + + def __init__(self, size: int, rank: int): + self._size, self._rank = size, rank + + def size(self, dim: int = 0) -> int: + return self._size + + def get_local_rank(self, dim: int = 0) -> int: + return self._rank + + +# --------------------------------------------------------------------------- +# Config translation +# --------------------------------------------------------------------------- + + +def test_domain_parallel_policy_passes_through_except_domain_size(): + """Everything in ``domain_parallelism`` but ``domain_size`` reaches the readers + unchanged, so their own validation sees what the user wrote (typos included).""" + cfg = OmegaConf.create( + { + "domain_parallelism": { + "domain_size": 4, + "auto_shard_size": 2048, + "placements": {"boundaries.stl_geometry": "replicate"}, + "typo_key": 1, + } + } + ) + policy = _domain_parallel_from_cfg(cfg) + assert policy == { + "auto_shard_size": 2048, + "placements": {"boundaries.stl_geometry": "replicate"}, + "typo_key": 1, + } + assert isinstance(policy["placements"], dict) # plain, not DictConfig + + # Unset block: an empty policy (readers apply their defaults). + assert _domain_parallel_from_cfg(OmegaConf.create({})) == {} + + +# --------------------------------------------------------------------------- +# Samplers over the data-parallel axis +# --------------------------------------------------------------------------- + + +def test_directory_samplers_shard_over_data_mesh_only(): + """Every rank of a domain group (same ddp rank) gets the same indices.""" + train_ds = list(range(40)) + val_ds = list(range(12)) + seqs = {} + for ddp_rank in (0, 1): + train_sampler, val_sampler = _build_directory_samplers( + train_ds, + val_ds, + use_distributed=True, + sampler_seed=7, + data_mesh=_FakeMesh(size=2, rank=ddp_rank), + ) + train_sampler.set_epoch(0) + seqs[ddp_rank] = (list(train_sampler), list(val_sampler)) + # Two ddp ranks partition the data ... + assert not set(seqs[0][0]) & set(seqs[1][0]) + assert sorted(seqs[0][1] + seqs[1][1]) == val_ds + # ... and are deterministic, so two domain ranks with the same ddp rank + # (which build the very same sampler) agree. + again, _ = _build_directory_samplers( + train_ds, + val_ds, + use_distributed=True, + sampler_seed=7, + data_mesh=_FakeMesh(size=2, rank=0), + ) + again.set_epoch(0) + assert list(again) == seqs[0][0] + + +def test_manifest_samplers_shard_over_data_mesh_only(): + """Manifest indices are split by ddp rank, not by world rank.""" + train_idx = list(range(30)) + val_idx = list(range(30, 40)) + parts = [] + for ddp_rank in (0, 1, 2): + train_sampler, val_sampler = _build_manifest_samplers( + train_idx, + val_idx, + dist_manager=None, # unused when a data_mesh is given + sampler_seed=3, + data_mesh=_FakeMesh(size=3, rank=ddp_rank), + ) + parts.append((set(train_sampler), set(val_sampler))) + assert set().union(*(p[1] for p in parts)) == set(val_idx) + for a in range(3): + for b in range(a + 1, 3): + assert not parts[a][0] & parts[b][0] diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel_gradients.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel_gradients.py new file mode 100644 index 0000000000..59203f7473 --- /dev/null +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/tests/test_domain_parallel_gradients.py @@ -0,0 +1,159 @@ +# 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. + +"""Gradient equivalence of the recipe's training step under domain parallelism. + +``train.py`` syncs the model over the domain mesh and then wraps it in DDP over +the flat world, relying on every parameter gradient arriving at DDP as the +complete domain-group value (ShardTensor resolves the sharded-axis reductions +before ``.grad`` accumulates). This test runs the recipe's ``forward_pass`` + +``_reduce_and_average`` on one sample sharded over the whole world and checks +that the loss and every parameter gradient equal the unsharded single-process +result. Skipped unless launched with more than one process:: + + torchrun --nproc-per-node 2 -m pytest tests/test_domain_parallel_gradients.py +""" + +from __future__ import annotations + +import pytest +import torch +from torch.distributed.tensor.placement_types import Shard + +pytest.importorskip("tensorboard") + +from collate import build_collate_fn # noqa: E402 +from loss import LossCalculator # noqa: E402 +from metrics import MetricCalculator # noqa: E402 +from train import _reduce_and_average, forward_pass # noqa: E402 + +from physicsnemo.distributed import DistributedManager # noqa: E402 +from physicsnemo.domain_parallel import ( # noqa: E402 + scatter_tensor, + sync_module_over_mesh, +) +from physicsnemo.mesh import DomainMesh, Mesh # noqa: E402 + +pytestmark = [pytest.mark.timeout(180)] + + +@pytest.fixture(scope="module") +def distributed_mesh(): + """1-D mesh over every launched process; skips outside a multi-process launch.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + DistributedManager.initialize() + dm = DistributedManager() + if dm.world_size < 2: + pytest.skip("launch with torchrun --nproc-per-node >= 2") + yield dm.initialize_mesh([-1], ["domain"]) + + +_N_POINTS = 1001 # uneven on 2/4/8 ranks +_TARGETS = {"pressure": "scalar", "wss": "vector"} + + +class _PointMLP(torch.nn.Module): + """Point-wise model: coordinates -> 4 outputs (pressure + wss).""" + + def __init__(self) -> None: + super().__init__() + self.net = torch.nn.Sequential( + torch.nn.Linear(3, 32), torch.nn.GELU(), torch.nn.Linear(32, 4) + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.net(x) + + +def _sample(device) -> DomainMesh: + torch.manual_seed(11) + points = torch.randn(_N_POINTS, 3, device=device) + interior = Mesh( + points=points, + point_data={ + "pressure": torch.randn(_N_POINTS, device=device), + "wss": torch.randn(_N_POINTS, 3, device=device), + }, + ) + return DomainMesh(interior=interior, boundaries={}, global_data={}) + + +def _shard_sample(domain: DomainMesh, device_mesh) -> DomainMesh: + """The same sample with interior points and point_data as Shard(0).""" + interior = domain.interior + shard = lambda t: scatter_tensor(t, 0, device_mesh, (Shard(0),)) # noqa: E731 + sharded = Mesh( + points=shard(interior.points), + point_data={k: shard(v) for k, v in interior.point_data.items()}, + ) + return DomainMesh(interior=sharded, boundaries={}, global_data={}) + + +def _step(model, domain, dist_manager): + """One recipe training step on *domain*: reduced loss and parameter grads.""" + collate = build_collate_fn("tensors", {"x": "interior.points"}, _TARGETS) + batch = collate([(domain, {})]) + loss_calc = LossCalculator(target_config=_TARGETS, loss_type="mse") + metric_calc = MetricCalculator(target_config=_TARGETS, metrics=["l2"]) + model.zero_grad(set_to_none=True) + loss, losses, metrics = forward_pass( + batch, + model, + "float32", + loss_calc, + metric_calc, + output_type="tensors", + target_config=_TARGETS, + ) + loss.backward() + avg_loss, _, _ = _reduce_and_average( + loss.detach(), losses, metrics, 1, device=dist_manager.device + ) + grads = [p.grad.detach().clone() for p in model.parameters()] + return avg_loss, grads + + +def test_domain_parallel_step_matches_single_process(distributed_mesh): + """Sharding the sample over the whole world (domain_size == world_size) + reproduces the unsharded loss and parameter gradients through the recipe's + sync -> DDP -> forward_pass -> reduce path.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA is not available") + + dm = DistributedManager() + device = dm.device + domain = _sample(device) + + # Reference: plain single-process step on the full sample. + torch.manual_seed(3) + reference = _PointMLP().to(device) + ref_loss, ref_grads = _step(reference, domain, dm) + + # Domain-parallel: same weights on every rank, DDP over the flat world + # (the recipe's exact wiring), sample sharded over the domain mesh. + torch.manual_seed(3) + model = _PointMLP().to(device) + sync_module_over_mesh(model, distributed_mesh) + ddp = torch.nn.parallel.DistributedDataParallel( + model, device_ids=[dm.local_rank], output_device=device + ) + dp_loss, dp_grads = _step(ddp, _shard_sample(domain, distributed_mesh), dm) + + assert abs(dp_loss - ref_loss) < 1e-5 * max(1.0, abs(ref_loss)) + for ref, got in zip(ref_grads, dp_grads): + assert type(got) is torch.Tensor, "DDP parameter received a distributed grad" + torch.testing.assert_close(got, ref, atol=1e-5, rtol=1e-4) From a39e4d4a7bc7ddfda9077392224343fd7b490f73 Mon Sep 17 00:00:00 2001 From: Corey Adams <6619961+coreyjadams@users.noreply.github.com> Date: Thu, 10 Sep 2026 18:43:44 +0000 Subject: [PATCH 2/2] Disable DDP optimizations for Shard tensor for these models for now; they cause strange crashes; to be optimized still --- .../unified_external_aero_recipe/src/train.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py index 99d2c2c02f..cbcede11d1 100644 --- a/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py +++ b/examples/cfd/external_aerodynamics/unified_external_aero_recipe/src/train.py @@ -949,6 +949,12 @@ def main(cfg: DictConfig) -> None: loaded_epoch = load_checkpoint(device=device, **ckpt_args) if cfg.compile: + if domain_mesh is not None: + # DDPOptimizer splits the compiled graph into gradient buckets and + # re-fakeifies tensors at each split; it does not understand + # ShardTensor and fails with "expected size ==". + # DDP itself is unaffected. + torch._dynamo.config.optimize_ddp = False model = torch.compile(model) num_epochs = cfg.training.num_epochs