Skip to main content

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

Project description

Jaxmg

JAXMg: A multi-GPU linear solver in JAX

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
  • cusolverMgPotri: Computes the inverse of an $N\times N$ symmetric (Hermitian) positive-definite matrix via a Cholesky decomposition.
  • cusolverMgSyevd: 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

Details for compiling the from source code can be found in CONTRIBUTING.md.

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), ))
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.

Projects that use JAXMg

  • JAXMg Benchmarks: Benchmarks for various Multi-GPUs setups.
  • JAXMg + Netket: Implementation of the MinSR Netket driver that uses JAXMg for inverting the S-matrix. Tested on Multi-node settings.
  • JAXMg for blurred sampling: Implementation of t-VMC that makes use JAXMg for inverting the QGT.

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

@misc{2601.14466,
Author = {Roeland Wiersema},
Title = {JAXMg: A multi-GPU linear solver in JAX},
Year = {2026},
Eprint = {arXiv:2601.14466},
}

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.7-cp314-cp314-manylinux_2_26_x86_64.whl (24.0 MB view details)

Uploaded CPython 3.14manylinux: glibc 2.26+ x86-64

jaxmg-0.0.7-cp313-cp313-manylinux_2_26_x86_64.whl (24.0 MB view details)

Uploaded CPython 3.13manylinux: glibc 2.26+ x86-64

jaxmg-0.0.7-cp312-cp312-manylinux_2_26_x86_64.whl (24.0 MB view details)

Uploaded CPython 3.12manylinux: glibc 2.26+ x86-64

jaxmg-0.0.7-cp311-cp311-manylinux_2_26_x86_64.whl (24.0 MB view details)

Uploaded CPython 3.11manylinux: glibc 2.26+ x86-64

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.7-cp314-cp314-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 3384b65cfe1170372fd71df392980fe8ff34f5f9d9b9a8d1a9c8697ed56ed93c
MD5 38920a68a67f83cbd27ef1ab237e1f11
BLAKE2b-256 a6cf1967c038cecc91466644efcae59ed6eed4181c6c658b1cd6b89eeabde72b

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.7-cp313-cp313-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 84a1baa3ddd5f594d9639f0b5227692419ed3e1d59fe2606aa368e9862590d87
MD5 19249c8fce164d8966e787ae83c475b3
BLAKE2b-256 2e6e31b34e0cdb5e83d4f496edd36911d9db21c73877017717797bf464543996

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.7-cp312-cp312-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 5c97db4f5fe918098ba8c73e4e75b2db6a463b05bad463c08fd71035439763be
MD5 f5dc1cd1413e93f85e3b13c066ded329
BLAKE2b-256 39e1f477ceba498c13a693ab8f33875058b5c268d168a93f4a6bee98b2f20071

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.7-cp311-cp311-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 2689174b928a8ade228166770d3e8960ba4baef44921695e18b31ef16a4c95ee
MD5 c284a547edb0a6d7a9f65f406e3dbd4c
BLAKE2b-256 da25fce55c76ff0f00cd70d829b777821bcc131e3af8e254ff5213688be03fb1

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