examples: give each MPI rank its own JAX device for the subspace expansion - #41
Merged
Merged
Conversation
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.
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.
Summary
The subspace expansion in
run_sqd_enlarge_subspace_sbd.pyis JAX, and JAX has noMPI 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 documentedworkaround 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_gpusconvention the README already documents forSBD instead of inventing a second policy.
Why
local_device_idsis passed explicitlyThe driver does use
cluster_detection_method="mpi4py"— that part works, since therendezvous fields come from
comm.Get_rank()/Get_size(). What it does not let JAXinfer 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 putsevery rank in its own node —
local_size=1everywhere, andMPI.Win.Allocate_sharedover
COMM_WORLDis refused outright ("processes are not in the same shared memorydomain"), 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.idis a global id in multi-process JAX: eightprocesses with one device each report
cuda:0…cuda:7whichever card they are on.device.local_hardware_idis the physical index, and it read0on every rank.The change
sbd.get_device_id()— new public query, backed by aplanned_device_idbinding.bindings.cppalready computesrank % device_countinside each diag entry point anddiscards it; this returns it.
The driver passes it as
local_device_ids, and additionally:XLA_PYTHON_CLIENT_PREALLOCATE=falseso XLA takes memory on demand instead ofreserving 75% away from SBD. Without it, SBD's Davidson basis runs out of room at
large
--max_dimand fails insidetpb_diagwithstd::bad_alloc— a memory errorthat reads as SBD's fault but is caused by JAX's reservation.
jax.config.update("jax_platforms", "cpu").JAX_PLATFORMS=cpu, and sets neither variable if the caller has.--device autoon rank 0 and broadcasts it._check_cuda()shells out tonvidia-smiwith a 2 s timeout and measures ~6 s when 8 ranks probe at once, so itcould time out on some ranks and not others. A diverged
device_strsent the ranksthat resolved to a GPU into
jax.distributed.initialize()— aCOMM_WORLDrendezvous— 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 samenvidia-smiprobe and printed "none detected" on a fully populated node when it timedout; 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):
--devicegpu(Thrust)gpu-omp(OMP offload)cpuautogpu, all ranks agree)gpu+ presetJAX_PLATFORMS=cpudevice.local_hardware_id:0..7across theeight ranks on both GPU backends, unset on
cpu. This is the check that read aflat
0before.--max_dimcompleted with JAXand SBD sharing each card and headroom remaining.
JAX_PLATFORMS— the configuration thatpreviously failed.
run_sqd_sbd.pywith--device autoon 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:
get_device_idlocal_hardware_idCUDA_VISIBLE_DEVICES=0,1,2,3on every rank there, so the placement comes from thischange rather than the launcher. Measured with the real
apply_excitationskernel: theexpansion output reports
platform=gpuandSingleDeviceShardingon the assigned card.The driver also reproduces the H2O reference above across nodes.
Notes for review
run_sqd_sbd.pygets the same--device autobroadcast, in its own commit. That bugis pre-existing and unrelated to JAX — that driver has no
jax.distributed.initialize— but a divergeddevice_strleaves ranks on differentbackends (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.cppchanged, so all three backends need rebuilding.