Skip to content

Custom-operation callback order under MPI - #50

Merged
msperryucsd merged 1 commit into
mainfrom
mpi_callback_issue
Oct 1, 2026
Merged

msperryucsd merged 1 commit into
mainfrom
mpi_callback_issue

Conversation

@spcvanschie

Copy link
Copy Markdown
Contributor

Human-written tl;dr: JAX sometimes mixes information from different forward evaluations or reverse-mode derivative calculations when multiple MPI ranks are used with custom CSDL operations, because it tries to run multiple concurrent calculation chains, even when MPI collective operations are used. This then means that the MPI collective operations start picking up the wrong information from their ranks. The code itself executes fine, but the output numbers are wrong and don't match check_totals().

The proposed fix just constrains JAX to run its calculation chunks in a specific order, rather than giving it absolute freedom to run them in whatever order. This fix is only active when MPI parallelization is used. I did some brief performance testing, from that the fix doesn't seem to negatively affect anything.

The rest of the mostly AI-generated text here contains a summary, a more elaborate description of the bug itself, a minimal reproducible example, a description of the proposed fix and added tests, and a performance comparison with and without the fix.

Summary

When a csdl_alpha model runs under MPI with the JAX backend, custom operations
that call MPI collectives (Allreduce, allgather, ...) in compute or
compute_jacvec_product can return wrong values and gradients. The results
differ from rank to rank, and no error is raised. This happens as soon as the
compiled program has enough independent callbacks, for example a model with
several objectives and constraints. It also affects csdl_alpha's own MPI
operations (mpi_sum, index_splitter, sync_variable and the mpi_region
derivatives).

The proposed fix instead runs custom-operation callbacks as ordered JAX callbacks
whenever more than one MPI process is running. Single-process models are unaffected.
A per-class setting, ordered_callbacks, can override the default. Compile and
run times are unchanged with one process, and the same or 7–16 % faster on 4
MPI ranks (see Performance).

The bug

The JAX backend lowers every CustomExplicitOperation (and
CustomImplicitOperation) to jax.pure_callback. That covers both the forward
pass (CustomExplicitOperation.compute_jax) and the reverse pass
(CustomJacOperation.compute_jax). XLA treats a pure callback as free of side
effects. That allows it to:

  • run independent callbacks in any order;
  • run them concurrently on its CPU worker threads. The XLA:CPU thunk
    executor hands ready work to a second thread once enough of it is
    independent; with jax 0.11.2 this happened at 10 or more independent
    callbacks;
  • skip callbacks whose outputs are unused.

Independent callbacks are common in csdl models:

  • csdl.derivative(..., loop=True) (the JaxSimulator.compute_totals default)
    builds one reverse sweep per output. The VJP callbacks of different
    objectives and constraints are therefore independent.
  • Several custom operations applied to the same input are independent in the
    forward pass.

When these callbacks contain MPI collectives, two collectives can be in flight
on one communicator at once, or start in a different order on different ranks.
MPI then matches each rank's collective with the wrong partner:

  • equal message sizes: silently wrong, rank-dependent results;
  • different sizes: an abort with Message truncated.

The order is decided on each rank independently, so a per-rank lock around the
callback is not enough: it serializes the calls, but ranks can still take
them in different orders. The existing derivatives_kwargs={'concatenate_ofs': True}
workaround serializes the reverse sweeps only. Independent forward-pass
operations still overlap.

Reproducer (numpy and scipy only)

Each rank holds a sparse block A_r. The custom operation computes
y = (sum_r A_r) x with an Allreduce in compute and another in the VJP.
Sixteen objectives f_k = w_k · y give 16 independent reverse sweeps. Every
rank also builds the global matrix with allgather at setup, so it can check
its gradients without any collective.

# OMP_NUM_THREADS=1 mpirun -n 4 python repro.py
from mpi4py import MPI
import numpy as np
import scipy.sparse as sp
import csdl_alpha as csdl

comm = MPI.COMM_WORLD
M, K = 400, 16                 # vector length, number of objectives

class DistributedMatvec(csdl.CustomExplicitOperation):
    """y = (sum over ranks of A_r) x, with an Allreduce in compute and in the VJP."""
    # with this MR, `ordered_callbacks = False` here restores the old behavior

    def __init__(self, A_local):
        super().__init__()
        self.A = A_local

    def evaluate(self, x):
        self.declare_input("x", x)
        y = self.create_output("y", (M,))
        self.declare_derivative_parameters("y", "x")
        return y

    def compute(self, inputs, outputs):
        outputs["y"] = np.empty(M)
        comm.Allreduce(self.A @ inputs["x"], outputs["y"], op=MPI.SUM)

    def compute_jacvec_product(self, inputs, outputs, d_inputs, d_outputs, mode):
        d_inputs["x"] = np.empty(M)
        comm.Allreduce(np.ascontiguousarray(self.A.T @ d_outputs["y"]), d_inputs["x"], op=MPI.SUM)

# rank-local sparse block, and the exact global matrix (no collectives needed later)
A_local = sp.random(M, M, density=0.01 * (comm.rank + 1), format="csr",
                    random_state=np.random.default_rng(comm.rank))
A_global = sum(comm.allgather(A_local))
x0 = np.sin(0.37 * np.arange(M)) + 0.1
W = np.random.default_rng(7).standard_normal((K, M))   # same on every rank

rec = csdl.Recorder(inline=False)
rec.start()
x = csdl.Variable(name="x", value=x0)
y = DistributedMatvec(A_local).evaluate(x)
fs = [csdl.sum(y * W[k]) for k in range(K)]           # K independent reverse sweeps
rec.stop()

sim = csdl.experimental.JaxSimulator(rec, gpu=False, additional_inputs=[x], additional_outputs=fs)
sim.run()
totals = sim.compute_totals()
err = max(np.linalg.norm(totals[fs[k], x].ravel() - A_global.T @ W[k]) / np.linalg.norm(A_global.T @ W[k])
          for k in range(K))
print(f"rank {comm.rank}: max relative gradient error {err:.1e}", flush=True)

On main (a69686b), every run gives wrong gradients, with errors that differ
between ranks:

rank 0: max relative gradient error 1.7e+00
rank 1: max relative gradient error 1.8e+00
rank 2: max relative gradient error 1.6e+00
rank 3: max relative gradient error 1.6e+00

With this MR, the gradients are exact on every rank (2.9e-16).

The new test file, csdl_alpha/src/operations/custom/tests/test_ordered_callbacks.py,
contains this check, and runs it under mpirun when called directly:

OMP_NUM_THREADS=1 mpirun -n 4 python test_ordered_callbacks.py reverse|forward [unordered] [concat]
  • forward uses 16 independent operations instead of 16 objectives.
  • unordered restores the old behavior.
  • concat adds concatenate_ofs.

Maximum relative error over all ranks, values and gradients (4 ranks, 16
independent callbacks):

Variant Reverse sweeps Forward operations
main behavior (unordered) 2.3e+00 3.0e+00
main + concatenate_ofs (unordered concat) 6.6e-13 4.8e+00
this MR (default) 6.6e-13 1.3e-15

During development, a version of this script also traced the callbacks on each
rank:

  • With 9 or fewer independent callbacks, main was correct in every run.
  • With 10 or more, it was wrong in every run, with 2 callbacks running at once
    on every rank (4 at once with 32 callbacks).
  • With this MR, each rank runs one callback at a time, also with 32.

Environment: Python 3.12.13, jax/jaxlib 0.11.2, numpy 2.5.3, scipy 1.18.1,
mpi4py 4.1.2, MPICH 5.0.1, Linux, CPU backend.

The fix

  1. host_callback in backends/jax/utils.py calls host Python from JAX,
    either unordered (sequential_pure_callback, as before) or ordered
    (jax.experimental.io_callback(..., ordered=True)).

    • Ordered callbacks run one at a time, in program order. That order is the
      same on every rank, because every rank traces the same csdl graph.
    • They also work inside scan/fori_loop/while_loop/cond, which csdl
      loops lower to.
    • JAX cannot vmap ordered callbacks. Under vmap (vectorized loops,
      csdl.derivative(..., loop=False)), the callback therefore falls back to
      unordered, with a warning.
  2. CustomOperation.ordered_callbacks (src/operations/custom/custom.py)
    selects the kind of callback for both CustomExplicitOperation.compute_jax
    and CustomJacOperation.compute_jax. It therefore applies to the forward
    and reverse passes of explicit and implicit operations alike.

    • None (default): ordered when MPI runs on more than one process. The
      check only consults mpi4py if the program has already imported it, so
      it never initializes MPI.
    • True or False: always or never ordered.

    It can be set on the class or on an instance. csdl_alpha's own MPI
    operations (mpi_sum, index_splitter, sync_variable) need no change:
    with more than one process the default orders them, and with one process
    their collectives cannot be mismatched.

  3. create_jax_function only traces the operations the outputs depend on
    (backends/jax/graph_to_jax.py). XLA removes unused pure callbacks, but it
    must run ordered ones.

    • The JAX function used to contain every operation in the graph. A run()
      function compiled after compute_totals() (which records the reverse
      pass into the root graph) would therefore execute every VJP callback on
      each call.
    • The function now contains only the ancestors of the requested outputs.
      That is what XLA's dead-code elimination already did for pure operations,
      so values do not change.
    • An unused RandomOperation still splits the PRNG key, so the random
      operations that do run get the same values as before.

Behavior changes

  • Single process: no change in lowering. The only difference is that
    unneeded operations are no longer traced, which XLA would have removed
    anyway.
  • More than one MPI process:
    • Custom-operation callbacks no longer overlap with each other. Python
      callbacks hold the GIL for most of their work, so little overlap was
      possible. Operations that want concurrency can set
      ordered_callbacks = False.
    • Inside vmap, callbacks remain unordered and a warning says so.
  • Random operations: an unused random operation in a graph without a PRNG
    key no longer raises an error, because it is skipped.

Performance

Measured against main (a69686b), with the two versions run alternately. The
numbers are medians over 3 processes; steady-state times are the median of
20 (single process) or 10 (MPI) calls.

These tables were measured with a first version of this MR, which emitted the
same callbacks. The final, shorter code was re-checked on the 2000-operation
model, again alternating with main: run() first call 2.14–2.19 s versus
2.11–2.22 s, and compute_totals() first call 10.06–10.24 s versus
10.24–10.37 s. The machine was more heavily loaded during this check.

Single process. Callbacks are unordered here, so only change 3 applies.
Compile and run times are the same within noise (≤ 2 %):

Model run() first call run() per call compute_totals() first call compute_totals() per call
2000 built-in operations, 8 outputs 2.07 → 2.03 s 1.86 → 1.88 ms 9.60 → 9.44 s 13.8 → 13.8 ms
40 custom operations, 8 outputs 0.36 → 0.36 s 6.54 → 6.63 ms 0.149 → 0.149 s 15.5 → 15.8 ms
same, run() compiled after compute_totals() 0.083 → 0.064 s 6.69 → 6.75 ms 0.448 → 0.443 s 15.7 → 15.8 ms

The last row is the case change 3 targets: run() compiles 23 % faster
because the reverse-pass operations are no longer traced.

MPI, 4 ranks. K independent custom operations, each doing about 14 ms of
BLAS work (which releases the GIL), so overlapping callbacks could in
principle help main. Times are the slowest rank, in seconds:

K Allreduce in callbacks run() per call compute_totals() per call compute_totals() first call
8 no 0.115 → 0.115 0.226 → 0.208 0.302 → 0.302
8 yes 0.117 → 0.119 0.231 → 0.216 0.307 → 0.328
16 no 0.259 → 0.229 0.486 → 0.409 0.629 → 0.593
16 yes 0.251 → 0.233 0.491 → 0.428 0.603 → 0.599

Ordering never made a model slower. With 16 operations, where main runs two
callbacks at once, it was 7–16 % faster: running the callbacks at the same
time gained nothing on main. With 16 operations and Allreduce, main's
results are also wrong (see above).

In a larger application, a 2D discontinuous Galerkin flow solver with mesh
warping (check_totals on 4 ranks), the ordered and unordered versions took
the same time (124–127 s).

Changes (Check all boxes that apply)

  • Added tests
  • Added examples
  • Updated docs
  • Refactor code
  • Added functionality
  • Fixed bugs (Patch version increase)
  • Changed API (Major/Minor version increase)
  • New release (Major/Minor/Patch version increase)

@spcvanschie

Copy link
Copy Markdown
Contributor Author

To add to the above, this issue comes up with the following package versions:

  • Python: 3.12.13
  • Numpy: 2.5.3
  • Scipy: 1.18.1
  • mpi4py: 4.1.2 on MPICH 5.0.1
  • csdl_alpha: 0.1.0, commit a69686b
  • jax: 0.11.2

@msperryucsd
msperryucsd merged commit b3cee4a into main Oct 1, 2026
2 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants