Skip to content

examples: give each MPI rank its own JAX device for the subspace expansion - #41

Merged
Sophia Wen (hfwen0502) merged 9 commits into
mainfrom
jax-per-rank-device
Sep 24, 2026
Merged

Sophia Wen (hfwen0502) merged 9 commits into
mainfrom
jax-per-rank-device

Conversation

@hfwen0502

@hfwen0502 Sophia Wen (hfwen0502) commented Sep 23, 2026 •

Copy link
Copy Markdown
Member

Summary

The subspace expansion in run_sqd_enlarge_subspace_sbd.py is JAX, and JAX has no
MPI awareness: every rank picks the first device it can see — the same device for
all of them. Since XLA reserves 75% of a device on first use, one rank claims it and
the rest die with RESOURCE_EXHAUSTED: CUDA_ERROR_OUT_OF_MEMORY. The documented
workaround was JAX_PLATFORMS=cpu, which sidesteps the problem rather than fixing it.

This gives each rank its own device by asking the backend which one it will use, so JAX
follows the same gpu_id = rank % num_gpus convention the README already documents for
SBD instead of inventing a second policy.

Why local_device_ids is passed explicitly

The driver does use cluster_detection_method="mpi4py" — that part works, since the
rendezvous fields come from comm.Get_rank()/Get_size(). What it does not let JAX
infer is the local device, because JAX's mpi4py path derives the local rank from
COMM_WORLD.Split_type(MPI.COMM_TYPE_SHARED), and on MPICH 5.0.0 (ch4:ucx) that puts
every rank in its own node — local_size=1 everywhere, and MPI.Win.Allocate_shared
over COMM_WORLD is refused outright ("processes are not in the same shared memory
domain"
), so the locality table itself is wrong, not just the split. Left to JAX, every
rank got local_device_ids=[0] and ran on physical GPU 0.

That went unnoticed because Device.id is a global id in multi-process JAX: eight
processes with one device each report cuda:0…cuda:7 whichever card they are on.
device.local_hardware_id is the physical index, and it read 0 on every rank.

The change

sbd.get_device_id() — new public query, backed by a planned_device_id binding.
bindings.cpp already computes rank % device_count inside each diag entry point and
discards it; this returns it.

The driver passes it as local_device_ids, and additionally:

  • sets XLA_PYTHON_CLIENT_PREALLOCATE=false so XLA takes memory on demand instead of
    reserving 75% away from SBD. Without it, SBD's Davidson basis runs out of room at
    large --max_dim and fails inside tpb_diag with std::bad_alloc — a memory error
    that reads as SBD's fault but is caused by JAX's reservation.
  • When SBD solver is running on CPUs, force the subspace expansion to CPU via jax.config.update("jax_platforms", "cpu").
  • honours an explicit JAX_PLATFORMS=cpu, and sets neither variable if the caller has.
  • resolves --device auto on rank 0 and broadcasts it. _check_cuda() shells out to
    nvidia-smi with a 2 s timeout and measures ~6 s when 8 ranks probe at once, so it
    could time out on some ranks and not others. A diverged device_str sent the ranks
    that resolved to a GPU into jax.distributed.initialize() — a COMM_WORLD rendezvous
    — while the rest skipped it, hanging the job. Broadcasting makes agreement structural,
    and runs one probe instead of eight.

Also dropped: the startup hardware report. print_device_info() re-runs the same
nvidia-smi probe and printed "none detected" on a fully populated node when it timed
out; the per-rank device assignment is now reported instead.

Testing

8 ranks, one node (x86_64 / H100, MPICH) — bundled H2O pool, all five paths exit 0
and reproduce the same energies (round 1 -76.2359466308, round 2 -76.2421767512,
final subspace 1742 x 1742, stopping on closure at round 2):

--device
gpu (Thrust) ✓
gpu-omp (OMP offload) ✓
cpu ✓ (CPU fallback, warns per rank)
auto ✓ (resolves to gpu, all ranks agree)
gpu + preset JAX_PLATFORMS=cpu ✓ (honoured)
  • Physical assignment verified via device.local_hardware_id: 0..7 across the
    eight ranks on both GPU backends, unset on cpu. This is the check that read a
    flat 0 before.
  • Oversubscription: 16 ranks on 8 GPUs pairs correctly (rank r → r % 8).
  • Large subspace: a long multi-round run at a high --max_dim completed with JAX
    and SBD sharing each card and headroom remaining.
  • Runs with no launcher wrapper and no JAX_PLATFORMS — the configuration that
    previously failed.
  • run_sqd_sbd.py with --device auto on 8 ranks: -76.2359466308.

8 ranks, two nodes (aarch64 / GB200, NVHPC 26.1, MPICH, Slurm) — 4 ranks and 4 GPUs
per node. Each rank lands on a distinct physical GPU within its own node:

rank node get_device_id local_hardware_id
0-3 A 0, 1, 2, 3 0, 1, 2, 3
4-7 B 0, 1, 2, 3 0, 1, 2, 3

CUDA_VISIBLE_DEVICES=0,1,2,3 on every rank there, so the placement comes from this
change rather than the launcher. Measured with the real apply_excitations kernel: the
expansion output reports platform=gpu and SingleDeviceSharding on the assigned card.
The driver also reproduces the H2O reference above across nodes.

Notes for review

  • run_sqd_sbd.py gets the same --device auto broadcast, in its own commit. That bug
    is pre-existing and unrelated to JAX — that driver has no
    jax.distributed.initialize — but a diverged device_str leaves ranks on different
    backends (the cpu and gpu extension modules are separate builds) and they then call
    solve_sci_batch, which is collective. Happy to split it out if you'd rather.
  • bindings.cpp changed, so all three backends need rebuilding.

The probe in front of jax.distributed.initialize() could hang the run. It
gated two COMM_WORLD collectives -- Split_type, and the bcast inside
initialize()'s coordinator lookup -- on a per-rank capability check:

    n_gpus = get_device_info().get('gpu_count') or 0
    if n_gpus > 0:
        ...Split_type... / ...initialize()...

get_device_info() reports gpu_count only when DeviceConfig._check_cuda()
succeeds first, and that probe runs bare nvidia-smi under a 2 s timeout, with
any exception swallowed as 'no GPU'. Measured on this node with persistence
mode disabled: 0.77-1.34 s for one call, 6.08 s across 8 concurrent ranks. So
ranks disagreed nondeterministically, and whichever entered the collective
blocked on the ranks that had skipped it.

Call initialize() unconditionally instead, and take --device gpu at its word:
if the user asked for a GPU backend, assume GPUs are present. No probe, no
branch, so no divergence to deadlock on. local_device_ids is left to JAX,
which derives the node-local rank itself via Split_type(MPI.COMM_TYPE_SHARED)
-- dropping the modulo that would otherwise hand two ranks the same device
when ranks outnumber GPUs, in favour of failing visibly.

Also drops the hardware-detected report, which was noise and, worse, was
printing 'GPU: none detected' on an 8x H100 node for the same timeout reason.

Verified on 8 ranks with no helper script and no JAX_PLATFORMS -- the
configuration that previously hung -- H2O in 18 s, rc=0, each rank on its own
card: rank r -> CudaDevice(id=r) for r in 0..7, one local device each.
Energies unchanged (-76.2359466308 -> -76.2421767512, 1742x1742, stopping on
closure at round 2).
The per-rank JAX device assignment added earlier did not work. JAX's mpi4py
cluster detection derives local_device_ids from
COMM_WORLD.Split_type(MPI.COMM_TYPE_SHARED), and MPICH 5.0.0 (ch4:ucx) puts
every rank in its own node: Split_type reports local_size=1 everywhere, and
Win.Allocate_shared over COMM_WORLD is refused outright, so the locality
table -- not just the split -- is wrong. Every rank therefore got
local_device_ids=[0] and ran on physical GPU 0. This was not caught because
Device.id is a GLOBAL id in multi-process JAX: eight processes with one
device each report cuda:0..7 whichever card they are on. device.local_hardware_id
is the physical index, and it read 0 on every rank.

Rather than re-derive the local rank in Python -- which needs either a
launcher-specific env var or the broken Split_type -- ask the backend what it
will use. bindings.cpp already computes rank % device_count to pick a device,
inside each diag entry point, and throws it away. planned_device_id() returns
it, exposed as sbd.get_device_id():

  - pure query: selects nothing and creates no context, so it is safe to call
    before JAX initialises its own backend, which is the whole requirement
  - device count comes from cudaGetDeviceCount/hipGetDeviceCount/
    omp_get_num_devices, not from parsing nvidia-smi, so it cannot flake. The
    existing nvidia-smi probe takes 0.77-1.34 s alone and 6.08 s across 8
    concurrent ranks against a 2 s timeout, which is what made
    DeviceConfig._check_cuda() report "GPU: none detected" on an 8x H100 node
  - it is deliberately a second copy of the rank % count rule rather than a
    refactor of the three diag sites, so exposing the value cannot change how
    any existing path selects its device

The driver now passes it to jax.distributed.initialize(), falls back to
JAX_PLATFORMS=cpu when no device is available, and honours an explicit
JAX_PLATFORMS=cpu. README updated: the example no longer needs
JAX_PLATFORMS=cpu, and the note explains the real failure -- including that
JAX's reservation surfaces as SBD's std::bad_alloc inside tpb_diag, which
reads as SBD's fault.

NOTE: bindings.cpp changed, so all three backends must be rebuilt before
sbd.get_device_id() exists. Until then the driver fails at import. Not yet
run end to end.
The CPU fallback did not work. JAX_PLATFORMS is a jax.config option parsed at
`import jax`, unlike XLA_PYTHON_CLIENT_*, which the C++ client reads when it
creates the backend -- so assigning os.environ after the import is silently
ignored. Measured: env-var-after-import leaves jax.devices() returning all
eight CudaDevice, while jax.config.update('jax_platforms', 'cpu') gives
[CpuDevice(id=0)].

Effect was that --device cpu on a GPU node printed "forcing the SQD subspace
expansion onto CPU" and then aborted (SIGABRT, rc=134) with
CUDA_ERROR_OUT_OF_MEMORY while trying to reserve 59.38GiB: every rank fell
back to JAX's default device with no reservation cap, since
XLA_PYTHON_CLIENT_PREALLOCATE=false is only set on the GPU branch.

Verified on 8 ranks, H2O, all four paths -- --device gpu (Thrust),
--device gpu-omp (offload), --device cpu, and --device gpu with
JAX_PLATFORMS=cpu preset. All four now exit 0 and reproduce the reference
energies: round 1 -76.2359466308, round 2 -76.2421767512, final subspace
1742x1742, stopping on closure at round 2. Physical device assignment
confirmed separately via device.local_hardware_id, which reads 0..7 across
the eight ranks on both GPU backends and is unset on cpu -- the check that
was flat at 0 before sbd.get_device_id() existed.
_check_cuda() shells out to nvidia-smi with a 2 s timeout, and at 8
concurrent ranks it has been measured at ~6 s. Every rank ran it
independently, so it could time out on some and not others.

A diverged device_str is not just a cosmetic disagreement here: the ranks
that resolved to a GPU go on to call jax.distributed.initialize(), a
COMM_WORLD rendezvous, while the ranks that fell back to cpu get -1 from
get_device_id() and skip it. The job then hangs until that rendezvous
times out.

Resolve on rank 0 and broadcast, so agreement is structural rather than
dependent on probe timing. Also one probe instead of eight, which removes
the concurrency that made it slow to begin with.

Verified on 8 ranks / 8x H100 with --device auto: reproduces the H2O
reference (round 1 -76.2359466308, round 2 -76.2421767512, final subspace
1742 x 1742).
Pre-existing, and independent of the JAX change: this driver has no
jax.distributed.initialize, but it resolves --device auto with the same
per-rank _check_cuda() probe. A diverged device_str leaves ranks on
different backends -- the cpu and gpu extension modules are separate
builds -- and they then call solve_sci_batch, which is collective over the
comm.

Verified on 8 ranks / 8x H100 with --device auto: -76.2359466308.
It referred readers to planned_device_id's docstring for the distinction
between the communicator rank and the node-local rank, but that docstring
no longer discusses it, so the pointer went nowhere.

State the rule instead and cite examples/README.md, which already
documents `gpu_id = rank % num_gpus` as SBD's rank-to-device convention.
bindings.cpp: "a refactor of the three of them" was ambiguous and
undercounted. There are three diag entry points, but the rank % count
arithmetic appears in four places -- those three plus the offload pin they
share. Name them instead of counting.

examples/README.md: "Both are skipped if you set them yourself" was true
only of XLA_PYTHON_CLIENT_PREALLOCATE. The device assignment is not an
environment variable and always runs; note that it stays correct when the
caller pins one GPU per rank, since a single visible card makes the count 1
and every rank resolves to device 0.
@hfwen0502
Sophia Wen (hfwen0502) merged commit 0625592 into main Sep 24, 2026
11 checks passed
@hfwen0502
Sophia Wen (hfwen0502) deleted the jax-per-rank-device branch September 24, 2026 19:56
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.

1 participant