Repository navigation
Custom-operation callback order under MPI - #50
Merged
Merged
Conversation
Contributor
Author
|
To add to the above, this issue comes up with the following package versions:
|
msperryucsd
approved these changes
Oct 1, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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, ...) incomputeorcompute_jacvec_productcan return wrong values and gradients. The resultsdiffer 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_variableand thempi_regionderivatives).
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 andrun 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(andCustomImplicitOperation) tojax.pure_callback. That covers both the forwardpass (
CustomExplicitOperation.compute_jax) and the reverse pass(
CustomJacOperation.compute_jax). XLA treats a pure callback as free of sideeffects. That allows it to:
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;
Independent callbacks are common in csdl models:
csdl.derivative(..., loop=True)(theJaxSimulator.compute_totalsdefault)builds one reverse sweep per output. The VJP callbacks of different
objectives and constraints are therefore independent.
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:
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 computesy = (sum_r A_r) xwith anAllreduceincomputeand another in the VJP.Sixteen objectives
f_k = w_k · ygive 16 independent reverse sweeps. Everyrank also builds the global matrix with
allgatherat setup, so it can checkits gradients without any collective.
On
main(a69686b), every run gives wrong gradients, with errors that differbetween ranks:
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
mpirunwhen called directly:forwarduses 16 independent operations instead of 16 objectives.unorderedrestores the old behavior.concataddsconcatenate_ofs.Maximum relative error over all ranks, values and gradients (4 ranks, 16
independent callbacks):
mainbehavior (unordered)main+concatenate_ofs(unordered concat)During development, a version of this script also traced the callbacks on each
rank:
mainwas correct in every run.on every rank (4 at once with 32 callbacks).
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
host_callbackinbackends/jax/utils.pycalls host Python from JAX,either unordered (
sequential_pure_callback, as before) or ordered(
jax.experimental.io_callback(..., ordered=True)).same on every rank, because every rank traces the same csdl graph.
scan/fori_loop/while_loop/cond, which csdlloops lower to.
vmap(vectorized loops,csdl.derivative(..., loop=False)), the callback therefore falls back tounordered, with a warning.
CustomOperation.ordered_callbacks(src/operations/custom/custom.py)selects the kind of callback for both
CustomExplicitOperation.compute_jaxand
CustomJacOperation.compute_jax. It therefore applies to the forwardand reverse passes of explicit and implicit operations alike.
None(default): ordered when MPI runs on more than one process. Thecheck only consults
mpi4pyif the program has already imported it, soit never initializes MPI.
TrueorFalse: 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.
create_jax_functiononly traces the operations the outputs depend on(
backends/jax/graph_to_jax.py). XLA removes unused pure callbacks, but itmust run ordered ones.
run()function compiled after
compute_totals()(which records the reversepass into the root graph) would therefore execute every VJP callback on
each call.
That is what XLA's dead-code elimination already did for pure operations,
so values do not change.
RandomOperationstill splits the PRNG key, so the randomoperations that do run get the same values as before.
Behavior changes
unneeded operations are no longer traced, which XLA would have removed
anyway.
callbacks hold the GIL for most of their work, so little overlap was
possible. Operations that want concurrency can set
ordered_callbacks = False.vmap, callbacks remain unordered and a warning says so.key no longer raises an error, because it is skipped.
Performance
Measured against
main(a69686b), with the two versions run alternately. Thenumbers 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 versus2.11–2.22 s, and
compute_totals()first call 10.06–10.24 s versus10.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 %):
run()first callrun()per callcompute_totals()first callcompute_totals()per callrun()compiled aftercompute_totals()The last row is the case change 3 targets:
run()compiles 23 % fasterbecause 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:Allreducein callbacksrun()per callcompute_totals()per callcompute_totals()first callOrdering never made a model slower. With 16 operations, where
mainruns twocallbacks at once, it was 7–16 % faster: running the callbacks at the same
time gained nothing on
main. With 16 operations andAllreduce,main'sresults are also wrong (see above).
In a larger application, a 2D discontinuous Galerkin flow solver with mesh
warping (
check_totalson 4 ranks), the ordered and unordered versions tookthe same time (124–127 s).
Changes (Check all boxes that apply)