Skip to content

Domain decomposition

When a model does not fit on one GPU, ModelParallel splits it into tiles — one GPU per tile — and runs the same physics with per-timestep halo exchanges over NCCL. It is the model-parallel counterpart to shot-parallel DDP (which replicates the whole model per rank and splits the shot list):

what is split each rank holds use it when
shot parallel (DDP) the shot list the whole model the model fits on one GPU
model parallel (DD) the model, into tiles one tile, all shots the model does not fit

Both compose: MeshTopology(py, px, shot_groups=...) describes a shot_groups × (py × px) rank grid, where ranks that share a tile coordinate run different shots and their gradients are all-reduced automatically after the backward.

The gradient is shot-summed, but the record is not: each shot group holds a different shot, so gather_record assembles per group, on that group's root rank (shot_group * py * px). With the default shot_groups=1 that root is rank 0 and every other rank gets None.

Hands-on companions: notebook 25 · Domain decomposition, notebook 26 · Overthrust 3-D, and the runnable FWI scripts under examples/multi-gpu/torch/dd_fwi_*.py.

Two classes

from sweep.parallel import MeshTopology, ModelParallel, pad_to_mesh

topo = MeshTopology(py=1, px=4, shot_groups=1,
                    world_size=world, rank=rank)      # 2-D: py must be 1
prop = PropTorch(eq, shape=(nz, nx), dh=dh, dt=dt, nt=nt, abcn=abcn,
                 impl="c", device=dev, ...)           # the GLOBAL problem spec
ddp  = ModelParallel(prop, topo)                      # wrap; tiles are automatic

ModelParallel reads the global problem off the wrapped single-domain PropTorch and derives everything per rank: the tile slice, the cut-aware padding (a cut face carries only the stencil halo, no PML — absorbing boundaries live only on true domain edges), global→tile source and receiver remapping, and the per-tile boundary-saving ring. The gradient-memory configuration (storage, storage_dtype, BoundarySaving.tail_steps) is inherited from the wrapped prop's memory config, so every rank is consistent by construction — see Boundary storage under DD for which values the DD backward actually accepts.

Launch with one process per GPU:

torchrun --standalone --nproc-per-node=4 your_script.py

--standalone means this machine only. To spread the same tiles over several machines, drop it and give the ranks a meeting point instead — every node then runs an identical command and negotiates its own numbering:

torchrun --nnodes=3 --nproc-per-node=4 \
  --rdzv-backend=c10d --rdzv-endpoint=<first-node>:29500 --rdzv-id=dd1 \
  your_script.py                      # py * px must now be 3 * 4 = 12

Nothing in your script changes: ModelParallel never picks a device and NCCL does not care whether a tile's neighbour is on this machine or the next one. The cross-node gradient is bit-identical to the single-domain one, same as within a node. A ready SLURM batch script and the pitfalls (one task per node, not per GPU; ask for partial nodes) are in examples/multi-gpu/torch/README.md.

Forward and gradients — plain autograd

The forward returns this rank's differentiable tile record; backward() produces model gradients exactly like the single-domain path:

vp_tile = torch.tensor(vp_global[..., ddp.x0:ddp.x0 + ddp.nxp],
                       device=dev, requires_grad=True)      # tile leaf
rec_tile = ddp(wavelet, src_global, rec_global, models=[vp_tile])
loss = misfit(rec_tile, obs_tile)
loss.backward()                       # vp_tile.grad = this tile's gradient
full_rec = ddp.gather_record(rec_tile)   # the shot group's root assembles

Two equivalent leaf styles are in use:

  • Tile leaf (above): each rank keeps only its slice; gradients stay per-tile. Least memory, no gradient collective.
  • Global leaf (the dd_fwi_* example scripts): every rank holds the full physical model, passes models=[pad_to_mesh(vp, px=px)], and adds one dist.all_reduce(vp.grad) after the backward. A single-GPU script becomes multi-GPU with two marked lines.

Each rank's record carries only the receivers its tile owns — ddp.own_receiver_indices gives the global indices, and ownership is a partition, so per-rank misfits over those traces sum to the global misfit. Accumulate the misfit scalar in float64 if you compare against a single-GPU run: the per-rank partial sums add in a different order than one GPU's single sum, which shows up at fp32 rounding otherwise.

The two explicit performance switches

models=None reuses the prepared model. Passing models= triggers the per-call model setup: tile slicing, edge padding, and a model-halo NCCL collective. When the model has not changed since the previous call — every shot of an observed-data loop, a line search — pass models=None to skip it. This is explicit on purpose: the propagator never guesses whether an in-place optimizer step changed the model. Autograd calls must keep passing the leaf tensors (the graph attaches to what you pass this call), so inversion loops re-pass models=[...] every shot; the per-step wavefield halo exchange always runs regardless.

torch.no_grad() keeps the capture forward-only. The first gradient-capable call allocates the adjoint wavefields and the nt-scaled boundary ring; a forward under no_grad with non-grad models does not. Generate observed data and run QC forwards under no_grad, and keep a separate ModelParallel instance for pure-forward work — an instance that has once produced a gradient keeps its adjoint machinery for its lifetime.

Correctness

On fp32 gpu-direct boundaries the DD gradient is bit-identical to the single-domain gradient (test/test_dd_backward_two_tile*.py, test/dd_api_check.py), including with boundary tail truncation (test/test_dd_tail_two_tile.py). The exception is AcousticVRZ3D at spatial order 2 or 4 (the default): its single-GPU backward uses a fused gradient kernel that DD cannot use, so the two differ by ~1 ULP per cell unless SWEEP_VRZ_GRAD_SPLIT=1 puts the single-GPU run on the same split kernel. pad_to_mesh pads the split axes up to the tile multiple; a padded run is a (slightly) different discrete problem than an unpadded one, so compare like against like — the example scripts' --check mode does exactly that.

Scope and limits

  • Equations: those with stepped compiled kernels — Acoustic (2-D), Acoustic3D, AcousticVRZ3D, Elastic (2-D) and Elastic3D. Everything else is refused at construction with an error that names the equation, including the 2-D AcousticVRZ (stepped, but its backward has no coupling-exchange phases), AcousticVTI/AcousticVTI1st, AcousticTTI, ElasticTTI, ElasticVRR and ViscoAcoustic.
  • Cuts: x strips in 2-D (py=1); x/y tile grids in 3-D.
  • Free surface: top face only under DD (a cut face can never carry one).
  • Topography: not supported. ModelParallel refuses a propagator built with topography= (NotImplementedError): the tiles carry no surface, and boundary saving gives a wrong gradient under a per-column surface anyway.
  • Gradient memory: boundary saving only. Each tile reconstructs its forward wavefield from saved boundaries, so a wrapped prop built with memory=Full() or memory=Ckpt(...) is refused at construction.
  • Boundary storage and dtype: see the table below.
  • BoundarySaving.tail_steps: Acoustic 2-D/3-D (see Propagators); it composes with cpu staging.
  • RTM: take the image through the gradient path (notebook 08 shows how); there is no separate rtm() entry. Encoded supershots (a (nsrc, nt) wavelet) are supported: each tile keeps the rows of the sources it owns (test/dd_encoded_check.py).
  • SWEEP_DD_DISABLE_OVERLAP=1 forces the serial step-then-exchange path — the bit-exact reference for the comm/compute-overlap forward and a production escape hatch.

Boundary storage under DD

The per-tile ring inherits storage / storage_dtype from the wrapped prop, but not every value is wired for the DD backward — which Python drives one step per kernel call, unlike the monolithic loop the staged paths were built for. What is refused, is refused loudly at the first backward:

BoundarySaving.storage Acoustic 2-D / 3-D, VRZ 3-D Elastic 2-D / 3-D, Elastic APM
"gpu" (default, gpu-direct) yes yes
"cpu" (pinned-host staging) yes yes — it used to raise
"disk" no — raises no — raises

Those two columns are the whole DD-admissible set: admission is declared by the equation's cuda_layout (see check_dd_admission), and nothing else declares it today. The refusal was lifted in the shared staggered driver, so the rest of that family (DAS-mu, elastic TTI SG, elastic VR) is no longer blocked by this check — but each still has to declare DD admission before any of it is reachable under ModelParallel.

For the acoustic and VRZ equations, storage="cpu" also needs a real cut: a single-tile ModelParallel (world_size=1) refuses it, because that path reaches a reconstruction indexing that is not exercised by any multi-tile run. Use gpu-direct there, or a plain PropTorch backward, which supports cpu and disk staging as usual.

Every storage_dtype works with either storage. On fp32 and bf16 the cpu-staged gradient is bit-identical to the gpu-direct one; fp16 and int8 differ only within their own run-to-run quantisation floor, i.e. by no more than two runs of the same configuration differ from each other (test/dd_offload_check.py checks exactly that, on both counts).

Staging trades PCIe traffic for GPU memory. The DD backward is driven one step per kernel call from Python, so the boundary runtime — a copy stream beside the compute stream, events either way, ring slots, and a prefetch of the next chunk issued while the current one is still being consumed — used to be built and destroyed inside every call, and a prefetch never survived to the step it was meant for. It is now owned by a BoundarySession held on the Python side for the whole reverse loop, so the stream and its ring events outlive the call that created them.

What it costs, measured on 2×V100 with a 3-D elastic tile (test/dd_elastic_staged_check.py), every arm bit-exact against gpu-direct:

ring config backward vs gpu-direct peak boundary memory
gpu-direct (baseline) 0.85 s 1.00× 1.03 GB
cpu, transfer_interval=1, ring_buffers=1 2.28 s 2.68× 0.29 GB
cpu, transfer_interval=8, ring_buffers=2 3.19 s 3.75× 0.37 GB
cpu, transfer_interval=32, ring_buffers=4 8.84 s 10.38× 0.81 GB

So it is still an escape hatch rather than a default — 3.5× less boundary memory for 2.68× the backward — but it is a working one, and the knobs matter more than the switch: 32/4 is the worst point measured here, not the best. The inherited cpu default (transfer_interval=64, ring_buffers=1, pinned) was not measured. Start at 1/1 and raise only if profiling asks.

API details: sweep.parallel reference.