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

File metadata

File hashes

Hashes for jaxmg-0.0.8-cp314-cp314-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 1720ecc6f700fd1a3c4aed4dbc31837f970f672d1bf7aceebd2eb0c449d5b74b
MD5 a4dfac3182b32e21f7fc4449c9086f57
BLAKE2b-256 1bcca2b53e478412ccb08cae2a53bf5d7bf3f0df088575521c56d87dfd01bc4c

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.8-cp313-cp313-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 3090ad0feb2ca947c26dac640ce885aa37b41b994c7e5c1f8143403e1990002d
MD5 e765a02e3b58e81f37ff3834d1240601
BLAKE2b-256 7c945f49632a167f884832a6c373e37452e0a5e8ffb0188e6f085afd22684ccd

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.8-cp312-cp312-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 f3787cc3f5a257db1839ac888f633ce0a811ffc060b5e84c557e862246a662e0
MD5 a22a3cdc68c8b78dad416a315c67c641
BLAKE2b-256 83a3f7de6c19f8faabf454c3238395f4781f9fc09673607a967c73e7a8b128f8

See more details on using hashes here.

File details

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

File metadata

File hashes

Hashes for jaxmg-0.0.8-cp311-cp311-manylinux_2_26_x86_64.whl
Algorithm Hash digest
SHA256 5f4e341e380bac17841b1db96105b30e926387c37c539a091e63df3640b36a69
MD5 e822c0ea84f2043bd66cafe90f1da186
BLAKE2b-256 b1f4d0b0b6668c7bbab34f59043c2dd37369ff838342995b4ba761813c63d2b0

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