Skip to content

MPI adjoint diverges from serial since 4.8.0 (float literal loses F suffix in UnaryOp printing) #3038

Description

@charlandellon

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.

Activity

  1. mloubout commented on Oct 1, 2026

    @mloubout
    Contributor

    _mpi_repro.py is missing

  2. charlandellon commented on Oct 2, 2026

    @charlandellon
    Author

    Reproducer: https://gist.github.com/charlandellon/800f340f181b66ae25bb43ad95da044d

    mpi_repro.py is fully self-contained (the reduced model is embedded, no torch/external data). Run:

    # serial reference
    mpirun -n 1 python mpi_repro.py --nt 400 --save ref.npy
    # decomposed (prints rel_err of the global gradient vs ref)
    mpirun -n 2 python mpi_repro.py --nt 400 --ref ref.npy

    Results (rel_err of the domain-decomposed adjoint vs serial; forward is exact):

    Devito rel_err
    4.7.1 4.48e-05 (correct)
    4.8.0 2.49e-02 (bug)
    4.8.23 2.81e-02 (bug)
    4.8.23 + the fix below 6.7e-06

    Fix (restores the float32 F suffix in UnaryOp/Cast), in devito/ir/cgen/printer.py:

    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}'
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions