Skip to content

Parallel (domain decomposition)

Model-parallel / domain-decomposition entry points. Usage guide: Domain decomposition; hands-on notebooks 25 and 26.

Rank grid

sweep.parallel.MeshTopology

MeshTopology(py: int, px: int, shot_groups: int, world_size: int, rank: int)

Rank-grid layout for (P_shot, Py, Px) decomposition.

The total world_size is partitioned as::

world_size = shot_groups * (py * px)

with each rank's coordinate derived as::

shot_group = rank // (py * px)
yi         = (rank %  (py * px)) // px
xi         = rank %  px

2-D problems must use py == 1.

Attributes

py, px : int Tile-grid extents along the y (crossline) and x (lateral) axes. For 2-D, py must be 1. shot_groups : int Number of orthogonal shot-parallel groups; ranks differing only in shot_group share the same tile coordinate (yi, xi). world_size, rank : int The familiar torch.distributed quantities. Kept on the topology object so all derived quantities (coords, neighbours, tile extents) are pure data.

coord property

coord

(shot_group, yi, xi) tuple for this rank.

tile_rank property

tile_rank

This rank's index INSIDE its shot group, i.e. yi * px + xi.

tile_rank == 0 marks the group's root -- the rank a per-shot-group collective gathers to. Global rank 0 is the root of shot group 0.

tile_world_size property

tile_world_size

Number of ranks per shot group, i.e. py * px.

is_edge

is_edge(axis: str, side: str) -> bool

Is this rank on the low / high edge along axis?

A 1-tile axis (px == 1 or py == 1) is considered to be on both edges — that single tile owns the absorbing boundary on both sides.

local_extent

local_extent(global_shape: Tuple[int, ...]) -> Tuple[Tuple[int, ...], Tuple[int, ...]]

Return (local_shape, offsets) for this rank's tile.

global_shape is the SWEEP wavefield layout: (Nz, Nx) for 2-D and (Nz, Ny, Nx) for 3-D. The depth axis Nz is never split.

v1 requires uniform tiles: Nx must be a multiple of px and Ny (3-D) a multiple of py. Padding to that multiple is the caller's responsibility (see README §7).

Returns

local_shape : tuple Same ndim as global_shape; Nz unchanged, Ny / Nx replaced by per-tile sizes. offsets : tuple (0, [oy,] ox) — global index of this tile's origin. The leading 0 is the (un-split) depth offset.

neighbour_rank

neighbour_rank(axis: str, direction: int) -> Optional[int]

Rank of the neighbour along axis in direction.

Parameters

axis : {'x', 'y'} direction : {-1, +1}

Returns

rank or None None when the rank sits at the global boundary along that axis (no neighbour exists) or when axis='y' and py == 1.

rank_at

rank_at(shot_group: int, yi: int, xi: int) -> int

Inverse of :pyattr:coord — return the rank with these coords.

The DD propagator wrapper

sweep.parallel.dd_propagator.ModelParallel

ModelParallel(prop, mesh: MeshTopology)

Run a single-domain :class:PropTorch decomposed across GPUs (model parallel / domain decomposition) — a strategy wrapper in the spirit of torch.nn.parallel.DistributedDataParallel::

prop = PropTorch(eq, shape=(nz, nx), dh=dh, dt=dt, nt=nt, abcn=abcn, ...)
ddp  = ModelParallel(prop, mesh)        # mesh = MeshTopology(py, px, ...)
syn  = ddp(wavelet, sources, receivers, models=[vp])   # same call as prop
loss = 0.5 * (syn - obs).pow(2).sum(); loss.backward()  # autograd-transparent

prop specifies the GLOBAL problem (equation, grid spacing, nt, abcn, spatial order, source/receiver types, free surface, PML, B) — the propagator you would build to run on one GPU if the model fit. ModelParallel reads that spec and builds per-tile solvers with cut-aware padding internally: a model-parallel grid cannot reuse one global prop's symmetric pad, so prop is a config carrier, not the compute object (constructing it is cheap — the big buffers are allocated lazily at forward, which only the per-tile solvers do). mesh is the :class:MeshTopology (x-cut py=1; 3-D may add a y-cut py>1 — see :func:sweep.parallel.balanced_grid).

own_receiver_indices property

own_receiver_indices

Global receiver indices (caller's receiver order) whose traces this rank's tile record carries — set by the last forward's geometry, empty before any call. Ownership is a partition of the global receiver list across the tile grid, so per-rank misfits over these traces sum to the global misfit (see the dd_fwi examples' partition assert).

forward

forward(wavelet, sources_global, receivers_global, models)

Run the DD forward; return this rank's tile record.

The record is in the SAME layout a single-domain PropTorch returns, (B, nt, nrec, nfield) -- only this rank's receivers (:attr:own_receiver_indices).

models is a list of global or already-tiled physical arrays, OR None to REUSE the model already edge-padded and halo-exchanged by a previous forward. An FWI epoch fires many shots through the SAME model (only the source moves), so re-padding the runtime model and running the NCCL model-halo collective on every shot is pure waste; pass models on the first shot of the epoch and models=None for the rest to skip it. (This is explicit on purpose — the propagator never guesses whether an in-place optimiser step changed the model, which would risk silently running on a stale model.)

AUTOGRAD: if any model tensor requires_grad AND grad mode is on, the returned record is differentiable — loss.backward() on a misfit populates each model's .grad (== :meth:gradient of the residual), exactly like the single-domain PropTorch autograd path. Otherwise it stays on the forward-only stepped path; the explicit :meth:gradient adjoint API remains for manual control.

MEMORY: the forward-only path allocates neither the adjoint wavefields nor the nt-scaled boundary ring — the capture that binds them is deferred until a gradient is first asked for. So wrapping a modelling / observed-data / line-search forward in torch.no_grad() is not cosmetic: it is how you avoid paying for gradient machinery you never use. An instance promoted to gradient-capable keeps that machinery for its lifetime (use_boundary_saving is frozen at capture), so build a separate ModelParallel for pure-forward work if you want it to stay light.

gather_record

gather_record(tile_record)

Assemble this shot group's record on the group root (None elsewhere).

The gather is per shot group, not world-wide: with shot_groups > 1 every group runs a DIFFERENT shot through the same tile grid, so ranks sharing a tile coordinate carry the same global receiver indices with different shots' traces. The root is global rank shot_group * py * px, which is rank 0 for the single-group case the guides describe. See :func:sweep.parallel.gather_tile_records.

Layout-agnostic on purpose: nrec is axis -2 both in the raw CUDA record and in the canonical (B, nt, nrec, nfield) one, so this needs no change when :meth:forward hands back the latter.

Mesh padding helpers

sweep.parallel.pad_to_mesh

pad_to_mesh(model, mesh=None, *, py: int = 1, px: int = 1)

Pad the trailing split axes up to the mesh's tile multiple.

Parameters

model : torch.Tensor or numpy.ndarray (nz, nx), (nz, ny, nx), or the same with leading batch / parameter axes. Trailing axes are padded on the HIGH side only: the last axis to a multiple of px, the second-to-last to a multiple of py. The depth axis is never a split axis and is never padded. mesh : MeshTopology or ModelParallelMesh, optional Source of py / px. Mutually exclusive with the keywords.

Returns

Same type as model; the input object itself when no pad is needed. Padding replicates the edge (np.pad(mode="edge") semantics), so the invented cells add no impedance contrast, and for torch tensors the operation is differentiable with the replicate adjoint (pad-cell gradients sum back onto the edge cells).

Notes

THE PAD IS NOT FREE, and its cost grows with the run. It adds cells in front of the high-side PML, so a padded run answers a slightly different problem than the unpadded one. Measured, acoustic 3-D, gradient relL2 against the same problem run WITHOUT the pad:

  • 80x95x95, 700 steps, 1-cell pad -> 1.7e-5, and 98 % of the difference sits in the padded faces;
  • a production-size grid, thousands of steps, encoded sources, 1-cell pad -> 1.3e-3, spread over the WHOLE volume (every cell differs; median distance to the nearest tile face 9 cells).

So "the pad only perturbs the edge" holds for short runs and stops holding for production ones — given enough steps the edge perturbation traverses the model.

It also scales with how close the ACQUISITION sits to the padded face, and that dependence is steep. Same 2-D setup, only the source/receiver positions moved: with the shot a third of the way in, the gradient moves 1.7e-7; with the shot 2 cells from the padded edge and a receiver on the last physical column, it moves 6.5e-2 — five orders of magnitude, because the pad pushes the PML one cell further from a near-boundary source and the gradient is near-source dominated. A survey that runs right up to the high edge will feel the pad far more than the headline numbers suggest. The indices themselves stay valid either way (see below); this is physics, not bookkeeping. DD itself stays bit-exact against a single domain on the SAME padded problem; the numbers above are the pad's own cost, not DD's. Prefer a rank count whose factors divide the grid — balanced_grid picks such a mesh when one exists and warns when it cannot. Do not expect to always dodge it: real grids are not friendly numbers. On a seven-band field cascade NO rank count divided all seven, and two bands admitted none at all above one rank, their extents being prime. In a multi-band cascade the pad is therefore unavoidable, which is why it is documented and warned about rather than designed away.

Sources and receivers need NO adjustment, which is the whole reason the pad is high-side only and never touches z: cell (iz, iy, ix) of the original model is still (iz, iy, ix) after padding, so every integer coordinate still addresses the same physical cell, and the free surface / sea floor do not move. Verified rather than asserted — with the same source/receiver arrays used on a padded and an unpadded run, the trace furthest from the padded face is bit-identical (relL2 = 0). Pad on the LOW side instead and every coordinate would need remapping.

Anything else defined per cell must be padded the SAME way, though — water masks, seabed depths, gradient masks. This function pads bool arrays too, so pad_to_mesh(mask, mesh) is the answer; a mask left unpadded is a shape mismatch, which is loud, but a mask padded with a different rule is not.

WHERE in the chain you pad matters, and getting it wrong is silent. Pad the TENSOR you hand the solver — i.e. after any reparameterisation has rendered it. Padding the stored model instead (growing the .npy before it is read) looks equivalent and is not: an INR renders velocity on normalised coordinates, so changing nx moves every sample point and the network paints a DIFFERENT model on the columns you already had. Measured on a low-frequency field velocity_inr case, a one-column pad of the input file changed 87 % of the shared cells by hundreds of m/s — swamping the ~1e-3 boundary effect it was meant to isolate.

Two ways to use this differ in what happens to the pad's gradient, and they are NOT equivalent:

  • RECOMMENDED — keep the UNPADDED tensor as the optimisation variable and call this inside the loss closure. The pad cells stay tied to the edge cell and their gradient folds back onto it through the replicate adjoint (exact integer factors; see test_model_parallel_padding.py). Know the side effect: folding dx copies back onto one column gives that column a gradient weight of 1 + dx — 2x for a 1-cell pad, but 12x for px=12 at its worst case, and (1+dy)(1+dx) at the corner (16x for py4 x px4). It is the true gradient of this parameterisation, not an error, but it does mean the last physical column and corner update faster than their neighbours. Detaching the copies would remove the weighting and also make the gradient inconsistent with the loss actually evaluated, so it is deliberately not done here; if the spike matters, prefer a mesh with a smaller pad.
  • Pad once, optimise the PADDED array, strip it at the end with :func:unpad_from_mesh. Now the pad cells are independent parameters: they drift off the edge value as the inversion proceeds (so they stop being a replicate), and unpad_from_mesh is a plain slice, so whatever gradient they accumulated is discarded instead of returning to the edge column. On a production run the pad column carried ~8 % of the peak gradient magnitude, so this is not a rounding-level difference. Freeze the pad region if you take this route.

sweep.parallel.unpad_from_mesh

unpad_from_mesh(arr, orig_shape: Sequence[int], mesh=None, *, py: int = 1, px: int = 1)

Slice a padded array back to the physical grid — loudly.

orig_shape is the trailing physical grid shape (2 or 3 ints). The trailing shape of arr must equal exactly what :func:pad_to_mesh would produce for this mesh; anything else (a stale shape after resampling, the wrong mesh) raises instead of silently mis-cropping. Leading batch axes pass through unchanged.

Building blocks

ModelParallel composes these; they are public so a caller can drive a mesh by hand (a custom halo pattern, a static partition computed ahead of time) without reimplementing the arithmetic.

sweep.parallel.mesh.ModelParallelMesh

ModelParallelMesh(
    grid: Tuple[int, int],
    shot_groups: Optional[int] = None,
    *,
    world_size: Optional[int] = None,
    rank: Optional[int] = None
)

Topology + ProcessGroups for (P_shot, Py, Px) decomposition.

Two sub-groups are created at construction time (one new_group call per ranks-list — every rank participates in every new_group call, even ones it isn't a member of, as required by PyTorch's collective contract):

model_pg The py * px tile ranks of THIS rank's shot group. Used by the halo exchange and the model-parallel gradient gather.

shot_pg The shot_groups ranks holding the same (yi, xi) tile across different shot groups. Used by the existing shot-parallel all_reduce on vp.grad.

Parameters

grid : (Py, Px) Tile-grid extents. Py == 1 for 2-D. shot_groups : int, optional Number of shot-parallel groups. Defaults to world_size // (Py*Px). world_size, rank : int, optional Override the values pulled from torch.distributed. Useful only for testing or non-default process-group setups.

sweep.parallel._topology.balanced_grid

balanced_grid(
    world_size: int,
    global_shape: Tuple[int, ...],
    *,
    shot_groups: int = 1,
    max_py: int = 2
) -> Tuple[int, int]

Recommend a (py, px) DD grid that keeps each rank's tile compact.

A 1-D x-cut (py=1, px=world) shrinks the contiguous x-dimension to Nx/world; on 8 GPUs that is the slow, memory-heavy choice. Measured on 8x V100 (acoustic 3-D, end-to-end forward via :class:DDPropagator), a balanced 2-D cut is substantially faster and lighter for strong scaling because cutting more axes saves more cut-aware PML and a squarer tile runs the FD kernel more efficiently:

=================== ============ ================== ======== global (Nz,Ny,Nx) x-cut px8 balanced px4,py2 speedup =================== ============ ================== ======== 256 x 256 x 1024 0.893 ms 0.814 ms +9 % 384 x 384 x 384 1.026 ms 0.716 ms +43 % 512 x 512 x 512 1.890 ms 1.439 ms +31 % =================== ============ ================== ========

(peak memory also drops ~17-19 %.) The win grows the thinner the x-cut tile would be — i.e. for cubic / strong-scaling problems.

This returns the (py, px) factorisation of the per-shot-group tile count (world_size // shot_groups) that MAXIMISES the smaller horizontal tile edge min(Ny/py, Nx/px), breaking ties toward a fatter (contiguous) x edge. 2-D models always get (1, px) (no y to split).

max_py caps the y-tile count (default 2 — a conservative load-balance choice that already captures most of the win). Raise it (e.g. max_py=world_size) to also consider py >= 3 — the fastest option for cubic globals (e.g. 384^3 -> px2py4 = 0.688 ms / 6.96x), fully validated (8x V100, bit-exact). [A past py>=3 boundary-save crash was a DDPropagator._capture cut_face_mask=0 bug, fixed pure-Python — not a kernel limitation.]

sweep.parallel.routing.partition_global_coords

partition_global_coords(
    coords_global: Tensor, topology: MeshTopology, global_shape: Tuple[int, ...]
) -> Tuple[torch.Tensor, torch.Tensor]

Select and shift the coords that fall inside THIS rank's tile.

Parameters

coords_global : torch.Tensor Shape (nshots, npts, ndim) integer tensor of GLOBAL indices. ndim must equal len(global_shape) (2 or 3). topology : MeshTopology Mesh layout; only the local (yi, xi) offset is consulted. global_shape : tuple of int (Nz, [Ny,] Nx) — used to compute the tile extent.

Returns

local_coords : torch.Tensor Same shape as coords_global. For on-tile entries the split axes are shifted to local indexing (x - ox, y - oy); off-tile entries are zeroed so they cannot be silently used as valid indices. mask : torch.Tensor Shape (nshots, npts) boolean tensor, True where the coord is inside this rank's tile.

sweep.parallel.routing.gather_tile_records

gather_tile_records(
    tile_record: Tensor, own_rec_idx: Sequence[int], mesh: Any
) -> Optional[torch.Tensor]

Reassemble one shot group's record from its tiles, on the group root.

The inverse of :func:partition_global_coords: that call handed each tile the receiver columns whose global coordinates fall inside it, this one puts those columns back at their global index.

The gather runs over mesh.model_pg -- the py * px ranks that decompose ONE shot -- and NOT over the world. With shot_groups > 1 every group propagates a DIFFERENT shot through the SAME tile grid, so two ranks sharing a tile coordinate carry the same global receiver indices holding different shots' traces. A world-wide gather writes both into one array and whichever tile is assembled later silently wins, which is a corrupt record rather than an error.

Parameters

tile_record This rank's record, (..., nrec_tile, nt). own_rec_idx Global receiver indices of this tile's columns, in tile order. mesh A :class:sweep.parallel.ModelParallelMesh.

Returns

torch.Tensor or None The assembled record on the shot group's root rank (topology.tile_rank == 0, i.e. global rank shot_group * py * px), None on every other rank.

sweep.parallel.routing.assemble_tile_records

assemble_tile_records(
    gathered: Sequence[Tuple[Sequence[int], Tensor]],
) -> Optional[torch.Tensor]

Place each tile's receiver columns at their global index.

Pure and collective-free so it can be tested without torch.distributed. Tiles that own no receiver contribute nothing; a global index no tile owns stays zero.

sweep.parallel.pml.build_rank_pml_widths

build_rank_pml_widths(
    mesh: ModelParallelMesh, abcn: int, ndim: int, *, image_method_active: bool = False
) -> List[int]

Per-side PML widths for THIS rank, in SWEEP's [z_low, z_high, (y_low, y_high,) x_low, x_high] layout.

Parameters

mesh : ModelParallelMesh Mesh topology (we only consult is_edge). abcn : int Global PML width (the value PML would have on every side if the rank were running un-split). ndim : int 2 or 3. image_method_active : bool If True, the z-low (top) PML is suppressed (the image method handles the free surface). Mirrors the existing single-rank behaviour in PropBase.init_abc.

Returns

list of int Length 2 * ndim. Pass directly to set_cpml_profiles_{s,r}(pml_width=...). Sides that don't have a neighbour-facing constraint get abcn; sides that face a neighbour get 0 (no PML).

Notes

A 1-tile axis (px==1 or py==1) is treated as having both edges, so PML is applied on both ends — matches the single-rank behaviour for that axis. :meth:MeshTopology.is_edge already returns True for both low and high in that case.

sweep.parallel.halo.HaloExchange

Bases: torch.autograd.function.Function

Exchange halo cells between neighbour ranks for one wavefield tensor.

Parameters

wavefield : torch.Tensor Shape (B, C, Nz, [Ny,] Nx_loc). Halo strips are written in place on the requested axes; the rest of the tensor is untouched. mesh : ModelParallelMesh Source of neighbour ranks and the model_pg process group. halo : int Halo width per side per axis (typically equation.so // 2). axes : tuple of {'x', 'y'} Split axes to exchange. 2-D: ('x',); 3-D: ('x', 'y').

Returns

wavefield : torch.Tensor The same tensor (halo regions updated in place).

Notes

Forward dispatches one fused dist.batch_isend_irecv per call covering all directions in axes (2-D: up to 2 messages; 3-D: up to 4). Backward runs the adjoint exchange (gradient halo → owner.interior += grad) plus zeros our own halo gradient where a neighbour exists.

sweep.parallel.halo.exchange_halos

exchange_halos(
    wavefields: Sequence[Tensor],
    mesh: ModelParallelMesh,
    halo: int,
    axes: Sequence[str],
) -> List[torch.Tensor]

Exchange halos for each field in wavefields via separate calls.

Convenience for multi-component equations. Each field calls HaloExchange.apply independently; messages are NOT fused across fields in this v1 (fusion is a follow-up optimisation when profiling flags per-step launch overhead).