Domain parallelism in the unified external aero recipe - #1983
coreyjadams wants to merge 2 commits into
Conversation
CODEOWNERS review mapCurrent for commit ⏳ @coreyjadams — 7 file(s)
⏳ @peterdsharpe — 7 file(s)
No CODEOWNER
Comment |
|
The implementation appears functionally sound, but the explicit model implementation standard must be satisfied before merging. FindingsSummary
Reviews (1) · Last reviewed commit: "Disable DDP optimizations for Shard tens..." |
| domain_size = int(cfg.get("domain_parallelism", {}).get("domain_size", 1)) | ||
| if domain_size <= 1: | ||
| return None, None |
There was a problem hiding this comment.
Invalid Sizes Silently Disable
domain_size is a GPU count, but zero and negative values are treated as if domain parallelism were disabled. A configuration such as domain_size=0 or domain_size=-1 therefore launches an ordinary unsharded run instead of rejecting the invalid topology. Validate that the value is positive before handling 1 as the disabled case.
| domain_size = int(cfg.get("domain_parallelism", {}).get("domain_size", 1)) | |
| if domain_size <= 1: | |
| return None, None | |
| domain_size = int(cfg.get("domain_parallelism", {}).get("domain_size", 1)) | |
| if domain_size < 1: | |
| raise ValueError("domain_parallelism.domain_size must be at least 1") | |
| if domain_size == 1: | |
| return None, None |
| _TARGETS = {"pressure": "scalar", "wss": "vector"} | ||
|
|
There was a problem hiding this comment.
The new _PointMLP directly subclasses torch.nn.Module, uses a non-raw summary-only class docstring, and annotates forward with plain torch.Tensor. The repository's model implementation standard requires model classes to inherit from physicsnemo.Module, use a raw docstring with Parameters, Forward, and Outputs sections, and use jaxtyping for public tensor annotations. This repository requirement must be satisfied before merging.
File Used: CODING_STANDARDS/MODELS_IMPLEMENTATION.md (source)
Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!
| 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 |
There was a problem hiding this comment.
7b96eb4 to
2280a38
Compare
2280a38 to
5d89fdb
Compare
`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.
…they cause strange crashes; to be optimized still
5d89fdb to
a39e4d4
Compare
| 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 | ||
|
|
||
|
|
There was a problem hiding this comment.
We use this for logs, printouts, to make sure we have full local tensors when needed.
| 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 <global>==<local>". | ||
| # DDP itself is unaffected. | ||
| torch._dynamo.config.optimize_ddp = False |
There was a problem hiding this comment.
This is problematic for performance still. Optimization is TBD though I don't want to block functionality: we need to address this to remove this, but it's an open still.
PhysicsNeMo Pull Request
domain_parallelism.domain_size=Nshards everysample over N GPUs; the remaining ranks form the data-parallel axis. Samples
arrive from the readers already sharded (previous PR) and the model runs on
ShardTensors, with DDP over the flat world for the parameter gradients.
domain_size=1leaves every code path identical to a non-domain-parallelrun.
conf/base.yaml:domain_parallelismblock (domain_size,auto_shard_size,placements); everything butdomain_sizeis thereaders'
domain_parallelpolicy and is passed through untouched.utils.build_distributed_meshes: builds the("ddp", "domain")devicemesh and returns
(domain_mesh, data_mesh), bothNonewhen off.datasets.build_dataloaders: keyword-onlydomain_mesh/data_mesh;readers get the mesh and 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;sync_module_over_meshbefore DDP;
materializeresolves sharded 0-D losses and metrics beforelogging. With
compile=true, DDPOptimizer is disabled under domainparallelism (it splits the graph at gradient buckets and does not
understand ShardTensor).
comparing loss and parameter gradients of a domain-parallel step against a
single process.
Validated on 4 GPUs with DrivAerML volume at 1M points: GeoTransolver and
FLARE train in bf16 with
domain_size=4; GeoTransolver also runs withcompile=true(currently slower than eager under domain parallelism; leftoff by default).
There are performance optimizations to do, here. Nevertheless, I want to get this open for review and feedback on the design while I tinker and get the speed fixed.
Description
Checklist
Dependencies
Review Process
All PRs are reviewed by the PhysicsNeMo team before merging.
Depending on which files are changed, GitHub may automatically assign a maintainer for review.
We are also testing AI-based code review tools (e.g., Greptile), which may add automated comments with a confidence score.
This score reflects the AI’s assessment of merge readiness and is not a qualitative judgment of your work, nor is
it an indication that the PR will be accepted / rejected.
AI-generated feedback should be reviewed critically for usefulness.
You are not required to respond to every AI comment, but they are intended to help both authors and reviewers.
Please react to Greptile comments with 👍 or 👎 to provide feedback on their accuracy.
Stack created with GitHub Stacks CLI • Give Feedback 💬