Parallel (domain decomposition)¶
Model-parallel / domain-decomposition entry points. Usage guide: Domain decomposition; hands-on notebooks 25 and 26.
Rank grid¶
sweep.parallel.MeshTopology ¶
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.
tile_rank
property
¶
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.
is_edge ¶
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 ¶
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 ¶
rank_at ¶
Inverse of :pyattr:coord — return the rank with these coords.
The DD propagator wrapper¶
sweep.parallel.dd_propagator.ModelParallel ¶
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
¶
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 ¶
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 ¶
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 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: foldingdxcopies back onto one column gives that column a gradient weight of1 + dx— 2x for a 1-cell pad, but 12x forpx=12at 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), andunpad_from_meshis 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 ¶
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).