Distributed Execution¤
JAXMg requires one Python process per GPU. It checks this requirement during a
solver call and raises a RuntimeError if the execution setup is unsupported.
This guide assumes two nodes with four GPUs per node, giving eight global devices:
node 0 node 1
------------------------- -------------------------
rank 0 -> local GPU 0 rank 4 -> local GPU 0
rank 1 -> local GPU 1 rank 5 -> local GPU 1
rank 2 -> local GPU 2 rank 6 -> local GPU 2
rank 3 -> local GPU 3 rank 7 -> local GPU 3
1. Initialize distributed JAX¤
Call jax.distributed.initialize()
before constructing the mesh or performing any JAX computation:
import jax
jax.distributed.initialize()
While JAX supports automatic detection of the distributed configuration, the user can provide the configuration explicitly:
jax.distributed.initialize(
coordinator_address="node-0-hostname:12345", # Rank 0 host and shared port.
num_processes=8, # Total number of Python processes.
process_id=process_id, # Unique global rank of this process.
local_device_ids=[local_device_id], # Local GPU assigned to this process.
)
2. Check the process and device view¤
After initialization, print the process and device information:
print("Process index:", jax.process_index())
print("Process count:", jax.process_count())
print("Local devices:", jax.local_devices())
print("Global devices:", jax.devices())
On rank 0, the output should have the following form:
Process index: 0
Process count: 8
Local devices: [CudaDevice(id=0)]
Global devices: [CudaDevice(id=0), ..., CudaDevice(id=7)]
3. Construct the process grid¤
JAXMg uses an ordinary two-dimensional JAX mesh. For eight ranks, valid grid shapes include \((8,1)\), \((4,2)\), \((2,4)\), and \((1,8)\). The following code constructs a \(4\times2\) grid:
from jax.sharding import PartitionSpec as P
process_rows = 4
process_cols = 2
mesh = jax.make_mesh((process_rows, process_cols), ("pr", "pc"))
jax.set_mesh(mesh)
matrix_specs = P("pr", "pc")
These JAX mesh dimensions are used throughout the backend and hence are inherited by cuSOLVERMp process-grid. The first
axis, pr, identifies process rows; the second, pc, identifies process columns.
Calling jax.set_mesh(mesh) makes this concrete mesh available while JAX traces
the explicitly sharded computation.
You can inspect the resultant process-rank mapping selected by JAX:
import numpy as np
rank_grid = np.asarray(
[[device.process_index for device in row] for row in mesh.devices]
)
if jax.process_index() == 0:
print(rank_grid)
For the row-major \(4\times2\) mesh used in this example, rank 0 prints
[[0 1]
[2 3]
[4 5]
[6 7]]
JAXMg accepts regular row-major and column-major rank mappings. For a \(4\times2\) grid these are
To request an explicit regular mapping, construct the JAX mesh from devices sorted by process rank:
from jax.sharding import Mesh
devices = np.asarray(
sorted(jax.devices(), key=lambda device: device.process_index),
dtype=object,
)
row_major_mesh = Mesh(
devices.reshape((process_rows, process_cols), order="C"),
("pr", "pc"),
)
column_major_mesh = Mesh(
devices.reshape((process_rows, process_cols), order="F"),
("pr", "pc"),
)
mesh = row_major_mesh # Or select column_major_mesh.
jax.set_mesh(mesh)
4. Shard the matrix and solve input¤
PartitionSpec describes how each array dimension maps onto the process-grid
axes. A matrix is normally sharded over both axes:
import jax.numpy as jnp
from jax.sharding import NamedSharding
N = 8192
a = jnp.eye(N, dtype=jnp.float32)
a_sharding = NamedSharding(mesh, P("pr", "pc"))
a = jax.device_put(a, a_sharding)
On the \(4\times2\) grid, each process initially owns an
\((N/4)\times(N/2)\) rectangular shard of a.
A vector solve input can be sharded over process rows and replicated across process columns:
b_vector = jnp.ones((N,), dtype=a.dtype)
b_vector = jax.device_put(b_vector, NamedSharding(mesh, P("pr")))
For a solve input with shape \(N\times\mathrm{NRHS}\), shard its rows and leave its columns replicated:
b_matrix = jnp.ones((N, 1), dtype=a.dtype)
b_matrix = jax.device_put(b_matrix, NamedSharding(mesh, P("pr", None)))
The public solver accepts either representation. JAXMg adds any routing or tile padding required for a narrow solve input and redistributes it internally for cuSOLVERMp.
With distributed execution configured, continue to Choose a tile size \(T_A\) before selecting a solver example.