PropJax¶
class PropJax(
equation,
shape,
source_type=[],
receiver_type=[],
abcn=50,
free_surface=False,
dh=10.0,
dt=0.002,
device=None,
backend=None,
memory=None,
use_ckpt=None,
ckpt_chunks=100,
pml_type=None,
scan_unroll=1,
)
Implementation:
src/sweep/propagator/jax.py
JAX propagator built around jax.lax.scan. The gradient-memory strategy is
chosen with memory=; with none given it is chunk-style rematerialization.
Note
PropJax shares the same solver concepts as the PyTorch backend, but its
runtime behavior is shaped by JAX transforms rather than Python-side loops.
Parameters¶
equation(equation instance): The equation instance to be stepped in JAX.shape(tuple[int, ...]): Physical model shape before absorbing boundaries are added. Use(nz, nx)in 2D and(nz, ny, nx)in 3D.source_type(list[str], optional): Wavefield names used for source injection. Defaults toequation.default_source_fields; names (or theirFieldSpecaliases) must be source-capable fields of the equation.receiver_type(list[str], optional): Wavefield names sampled at receiver locations. Defaults toequation.default_receiver_fields; names or aliases must be receiver-capable fields.abcn(int, optional): Absorbing boundary width.free_surface(bool, optional): Whether the top boundary is treated as a free surface. This affects internal coordinate offsets before source injection and receiver sampling.dh(floator sequence, optional): Grid spacing: a scalar, or one value per axis,(dz, dx)in 2D and(dz, dy, dx)in 3D.dt(float, optional): Time step in seconds.device(device or context, optional): Stored device/context argument (dev=is a deprecated alias). Actual JAX execution placement is still driven by JAX arrays and transforms.backend(str, optional):'jax'orNone; accepted for symmetry withPropTorch.memory(optional): The gradient-memory strategy, one ofFull(),BoundarySaving(...)orCkpt(...)fromsweep.propagator.options.Full()differentiates through the scan tape;Ckpt(chunks=...)is chunkedjax.checkpointrematerialization (mode='chunk'only);BoundarySaving(...)reconstructs the forward wavefield in reverse time from saved boundaries, with the ring kept on the device (storage='gpu', notail_steps).None(default) means checkpointing.use_ckpt(bool | None, optional): Legacy switch:Truerequests chunked rematerialization,Falsefull storage.None(default) leaves the choice tomemory=, else checkpointing.ckpt_chunks(int, optional): Chunk size, in time steps, when checkpointing.pml_type(str, optional): PML formulation.None(default) usesequation.default_pml_type(e.g.'cpmlr'forAcoustic,'cpmls'forElastic);'spml'is supported byAcoustic1stonly. A value outsideequation.supported_pmlraisesValueError; leave it unset.scan_unroll(int, optional):lax.scanunroll factor for the time loop. Small, launch-bound grids can gain from 2–4 (gradients bit-identical); large, bandwidth-bound grids lose a few percent, so the default is1.
Forward Parameters¶
wavelet,sources,receivers: see Runtime Shape Conventions. Modes A1, A2, and B (source encoding) are auto-detected from the input shapes; the legacysource_encoding=kwarg has been removed.models(list of arrays, optional): List of model arrays in the exact order required byequation.models.return_wavefield(bool, optional): IfTrue, returns an auxiliary wavefield output in addition to the recorded data.adj(bool, optional): Adjoint-style forward switch.
Return Value¶
- default:
record - if
return_wavefield=True:(record, snapshots)