Devito issue report: MPI adjoint diverges from serial (float literal loses F suffix in UnaryOp)
Summary
Since Devito 4.8.0, the MPI (domain-decomposed) adjoint/gradient diverges
from the serial result on a realistic 3D acoustic model, while the forward is
exact. On Devito 4.7.1 the adjoint is correct. git bisect (v4.7.1 → v4.8.0)
points to commit 32f6e13 ("symbolics: switch to sympy printer settings and
improve UnaryOp printing").
Root cause: CodePrinter._print_UnaryOp prints its base with the current
printer instead of ccode(...), so float literals inside UnaryOp/Cast
lose the F dtype suffix and are promoted to double, e.g.
(int)(floor(5.0e-2F*posx)) (4.7.1) becomes (int)(floor(5.0e-2*posx)) (4.8.x).
The forward output is unchanged, but the MPI adjoint accumulates a
decomposition-dependent error that grows with propagation time.
Environment
- Devito 4.8.23 (pip), Python 3.11, GCC/OpenMP.
- MPI: OpenMPI 5.0.5 + mpi4py 4.1.2,
DEVITO_MPI=basic, mpirun -n 1 vs -n 2.
- 3D acoustic,
space_order=4, nbl=25, bcs='damp', float32.
- Model: heterogeneous two-layer / real cropped model; 16–1600 receivers.
Reproducer
_mpi_repro.py (self-contained, no torch):
# reference (serial)
mpirun -n 1 python _mpi_repro.py --nt 400 --save ref.npy
# distributed
mpirun -n 2 python _mpi_repro.py --nt 400 --ref ref.npy
It builds a Model/AcquisitionGeometry/AcousticWaveSolver, runs a forward
with save=True, sets a smooth residual (the forward data) as receiver data,
runs solver.jacobian_adjoint(..., u=u_saved), reconstructs the global gradient
via Function.local_indices + Allreduce, and reports the relative L2 error
vs the serial reference.
Results
| Devito |
adjoint MPI, nt=400 |
adjoint MPI, nt=601 (serial ref) |
| 4.8.23 |
2.81e-02 |
— |
| 4.8.0 |
2.48e-02 |
— |
| 4.7.1 |
4.4e-05 (OK) |
3.95e-05 (OK) |
- Forward (data and saved wavefield
u[t]) matches serial to ~1e-6 even in 4.8.x.
- Checkpointed adjoint (pyrevolve) in 4.8.x diverges at long propagation (~6e-2
at nt=1120); non-checkpointed diverges too, so it is not a checkpointing
issue.
- Not caused by: halo scheme (
basic/overlap), halo width (space_order=2),
optimization level (opt=noop), number of checkpoints.
- The error appears once the wavefield interacts with the absorbing boundary
(step between nt=300 exact and nt=400 diverging).
Bisect
git bisect start
git bisect good v4.7.1
git bisect bad v4.8.0
git bisect run bash bisect_test.sh
# 32f6e136720356f50afd2444db4b22e47f660482 is the first bad commit
Diff of the generated Gradient operator between the bad commit and its parent,
for the receiver cast:
- int ii_rec_0 = (int)(floor(5.0e-2F*posx));
+ int ii_rec_0 = (int)(floor(5.0e-2*posx));
i.e. the F suffix (float32 literal) is lost.
Suggested fix
In devito/ir/cgen/printer.py, _print_UnaryOp should print its base with
ccode (which applies the correct dtype literal suffix) rather than with the
current printer:
def _print_UnaryOp(self, expr, op=None, parenthesize=False):
op = op or expr._op
base = ccode(expr.base) # was: self._print(expr.base)
if not q_leaf(expr.base) or parenthesize:
base = f'({base})'
return f'{op}{base}'
(An equivalent fix is to route UnaryOp/Cast through ccode, as verified by
restoring the 4.7.1-style __str__ that used ccode.)
Validation of the fix
Applying just the one-line change above to 4.8.23:
| case |
before |
after |
| adjoint MPI, nt=400 |
2.81e-02 |
7.29e-06 |
| adjoint MPI, nt=1120 (checkpointed) |
~6e-02 |
2.53e-05 |
The MPI adjoint then matches the serial result.
Devito issue report: MPI adjoint diverges from serial (float literal loses
Fsuffix in UnaryOp)Summary
Since Devito 4.8.0, the MPI (domain-decomposed) adjoint/gradient diverges
from the serial result on a realistic 3D acoustic model, while the forward is
exact. On Devito 4.7.1 the adjoint is correct.
git bisect(v4.7.1 → v4.8.0)points to commit 32f6e13 ("symbolics: switch to sympy printer settings and
improve UnaryOp printing").
Root cause:
CodePrinter._print_UnaryOpprints its base with the currentprinter instead of
ccode(...), so float literals insideUnaryOp/Castlose the
Fdtype suffix and are promoted todouble, e.g.(int)(floor(5.0e-2F*posx))(4.7.1) becomes(int)(floor(5.0e-2*posx))(4.8.x).The forward output is unchanged, but the MPI adjoint accumulates a
decomposition-dependent error that grows with propagation time.
Environment
DEVITO_MPI=basic,mpirun -n 1vs-n 2.space_order=4,nbl=25,bcs='damp',float32.Reproducer
_mpi_repro.py(self-contained, no torch):It builds a
Model/AcquisitionGeometry/AcousticWaveSolver, runs a forwardwith
save=True, sets a smooth residual (the forward data) as receiver data,runs
solver.jacobian_adjoint(..., u=u_saved), reconstructs the global gradientvia
Function.local_indices+Allreduce, and reports the relative L2 errorvs the serial reference.
Results
u[t]) matches serial to ~1e-6 even in 4.8.x.at nt=1120); non-checkpointed diverges too, so it is not a checkpointing
issue.
basic/overlap), halo width (space_order=2),optimization level (
opt=noop), number of checkpoints.(step between nt=300 exact and nt=400 diverging).
Bisect
Diff of the generated
Gradientoperator between the bad commit and its parent,for the receiver cast:
i.e. the
Fsuffix (float32 literal) is lost.Suggested fix
In
devito/ir/cgen/printer.py,_print_UnaryOpshould print its base withccode(which applies the correct dtype literal suffix) rather than with thecurrent printer:
(An equivalent fix is to route
UnaryOp/Castthroughccode, as verified byrestoring the 4.7.1-style
__str__that usedccode.)Validation of the fix
Applying just the one-line change above to 4.8.23:
The MPI adjoint then matches the serial result.