Skip to main content

JAXMg provides a C++ interface between JAX and cuSolverMg, NVIDIA's multi-GPU linear solver.

Project description

Jaxmg

JAXMg: A distributed linear solver in JAX with cuSolverMg

Docs Releases Build Status

JAXMg

JAXMg provides a C++ interface between JAX and cuSolverMg, NVIDIA’s multi-GPU linear solver. We provide a jittable API for the following routines.

  • cusolverMgPotrs: Solves the system of linear equations: $Ax=b$ where $A$ is an $N\times N$ symmetric (Hermitian) positive-definite matrix via a Cholesky decomposition
  • cusolverMgPotrs: Computes the inverse of an $N\times N$ symmetric (Hermitian) positive-definite matrix via a Cholesky decomposition.
  • cusolverMgPotrs: Computes eigenvalues and eigenvectors of an $N\times N$ symmetric (Hermitian) matrix.

For more details, see the API.

Installation

The package is available on PyPi and can be installed with

pip install jaxmg[cuda12]

This will install a GPU compatible version of JAX.

  1. pip install "jaxmg[cuda12]": Use CUDA 12 (only works for jax>=0.6.2).

  2. pip install "jaxmg[cuda12-local]": Use locally available CUDA 12 installation.

  3. pip install "jaxmg[cuda13]": Use CUDA 13 (only works for jax>=0.7.2).

  4. pip install "jaxmg[cuda13-local]": Use locally available CUDA 13 installation.

The provided binaries are compiled with

JAXMg CUDA cuDNN
cuda12,cuda12-local 12.8.0 9.17.1.4
cuda13,cuda13-local 13.0.0 9.17.1.4

Note: pip install jaxmg will install a CPU-only version of JAX. Since jaxmg is a GPU-only package you will receive a warning to install a GPU-compatible version of jax.

Example

A minimal example that runs the code is:

import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
from jax.sharding import PartitionSpec as P, NamedSharding
from jaxmg import potrs
print(f"Devices: {jax.devices()}")
# Assumes we have at least one GPU available
devices = jax.devices("gpu")
N = 12
T_A = 3
dtype = jnp.float64
# Create diagonal matrix and `b` all equal to one
A = jnp.diag(jnp.arange(N, dtype=dtype) + 1)
b = jnp.ones((N, 1), dtype=dtype)
ndev = len(devices)
# Make mesh and place data (rows sharded)
mesh = jax.make_mesh((ndev,), ("x",))
A = jax.device_put(A, NamedSharding(mesh, P("x", None)))
b = jax.device_put(b, NamedSharding(mesh, P(None, None)))
# Call potrs
out = potrs(A, b, T_A=T_A, mesh=mesh, in_specs=(P("x", None), P(None, None)))
print(out)
expected_out = 1.0 / (jnp.arange(N, dtype=dtype) + 1)
print(jnp.allclose(out.flatten(), expected_out))

which gives

[[1.        ]
 [0.5       ]
 [0.33333333]
 [0.25      ]
 [0.2       ]
 [0.16666667]
 [0.14285714]
 [0.125     ]
 [0.11111111]
 [0.1       ]
 [0.09090909]
 [0.08333333]]
True

as expected.

cuSolverMp

As of CUDA 13, there is a new distributed linear algebra library called cuSolverMp with similar capabilities as cuSolverMg, that does support multi-node computations as well as >16 devices. Given the similarities in syntax, it should be straightforward to eventually switch to this API. This will require sharding data into a cyclic 2D form and handling the solver orchestration with MPI.

Citations

@software{Wiersema_JAXMg_distributed_linear_2025,
author = {Wiersema, Roeland},
month = dec,
title = {{JAXMg: distributed linear solvers in JAX}},
url = {https://github.com/flatironinstitute/jaxmg},
version = {0.0.3},
year = {2025}
}

Acknowledgements

I acknowledge support from the Flatiron Institute. The Flatiron Institute is a division of the Simons Foundation.

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distributions

No source distribution files available for this release.See tutorial on generating distribution archives.

Built Distributions

If you're not sure about the file name format, learn more about wheel file names.

jaxmg-0.0.3-cp314-cp314-manylinux_2_26_x86_64.whl (16.0 MB view details)

Uploaded CPython 3.14manylinux: glibc 2.26+ x86-64

jaxmg-0.0.3-cp313-cp313-manylinux_2_26_x86_64.whl (16.0 MB view details)

Uploaded CPython 3.13manylinux: glibc 2.26+ x86-64

jaxmg-0.0.3-cp312-cp312-manylinux_2_26_x86_64.whl (16.0 MB view details)

Uploaded CPython 3.12manylinux: glibc 2.26+ x86-64

jaxmg-0.0.3-cp311-cp311-manylinux_2_26_x86_64.whl (16.0 MB view details)

Uploaded CPython 3.11manylinux: glibc 2.26+ x86-64

File details

Details for the file jaxmg-0.0.3-cp314-cp314-manylinux_2_26_x86_64.whl.

File metadata

File hashes

Hashes for jaxmg-0.0.3-cp314-cp314-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 fb7cbf121cbb28af7d1af43632c2d7760e36668be90fdeb6199f5b1f5e92d099
MD5 4ef3cfe84ed1f56775ae966ce62be174
BLAKE2b-256 006e9ce9b61a489cc8be1d2cb6d551e85246c695041e8d9ad39357f514bc35b6

See more details on using hashes here.

File details

Details for the file jaxmg-0.0.3-cp313-cp313-manylinux_2_26_x86_64.whl.

File metadata

File hashes

Hashes for jaxmg-0.0.3-cp313-cp313-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 21a172b2eb84e7b792f8cf806cad238dd45cb720f2e8bf6abb791b178fdae513
MD5 0baa31a99bb49a3cfd14c1a4df41a26b
BLAKE2b-256 92989658188b70351d89cb902a351c452c363299783fe707266c6ea5b729dc3a

See more details on using hashes here.

File details

Details for the file jaxmg-0.0.3-cp312-cp312-manylinux_2_26_x86_64.whl.

File metadata

File hashes

Hashes for jaxmg-0.0.3-cp312-cp312-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 e5463055f1f67c107ae93802abe82ea2f9b1491e53301c8dbfc0fc6bdad640d0
MD5 2b94bf6226cbc909c598f2a0723eff78
BLAKE2b-256 4fc8d02797a84758a935567c86d5b7c913d33934d7f46f89eaba8f29042bf747

See more details on using hashes here.

File details

Details for the file jaxmg-0.0.3-cp311-cp311-manylinux_2_26_x86_64.whl.

File metadata

File hashes

Hashes for jaxmg-0.0.3-cp311-cp311-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 72550da49b5da322125ea3ad6676d49800a3a0ef9daad00d6bbaaa06d4edb00e
MD5 a9538d91c34b9988d1c428d36679fba0
BLAKE2b-256 12c8bbd3c5032436f9cdd24cb307bbc9e8fe1eb7581184bcc26828f445e38768

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page