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
  • 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

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

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 SR-matrix. Tested on Multi-node settings.

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.6-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.6-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.6-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.6-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.6-cp314-cp314-manylinux_2_26_x86_64.whl.

File metadata

File hashes

Hashes for jaxmg-0.0.6-cp314-cp314-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 17fa4cafdd0bab72364f89be623c345adc9ee89efcf5a67276314c66183e2154
MD5 ab50a92129ce0636e2cd39bf2d50bbdd
BLAKE2b-256 4a3b710d8191be4a7e770d64f18d0a37d0a8312601513cfc493f9178669722e1

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.6-cp313-cp313-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 90bcf860cce5fa8753b5f89627363f5485acee89fe6dd51829d48345f267797e
MD5 0439fcb166cf5a599e3d0e7336aa93c2
BLAKE2b-256 26d21e4c16b48d993e31751cd976057c1b74219c5763d113183cdfc80d563419

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.6-cp312-cp312-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 14250477f2102121feb7dca6a79d89af1a716c198743225f026b27b31fe5eded
MD5 4ad30dea7836e5186654ee2c73c35e34
BLAKE2b-256 6c605db08b4bedceb2d356a6ad8e939f0de7f47aeeea73bc96d62e7c731d0577

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.6-cp311-cp311-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 2a4967b72aefd0488e3ea72edcb9ceba09e3191980b58269ffef1d4af570c1bf
MD5 3075fea104cb2c4ee36011557f18770b
BLAKE2b-256 802e629b486a6361e95f0ba145864763f8f0ccf2504ab23276de37143b79fae4

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