Skip to content

API Reference

This page documents the public API of JAX-AMG.

Solver

jaxamg.solve(A, b, x0=None, config=None, block_dim=1, comm=None, nglobal=None, partition_info=None, save_stats_file=None, reuse_setup=False, **kwargs)

Solve Ax=b using the AmgX backend. See Examples for usage.

Parameters:

Name Type Description Default
A MatrixOrOperator

Matrix or callable operator A(x). All matrices/operators are converted to jax.experimental.sparse.bcsr sparse matrices internally. In MPI mode this is the local partition.

required
b ArrayLike

Right-hand-side vector. In MPI mode this is the local RHS partition.

required
x0 ArrayLike | None

Optional initial guess (same shape as b; local partition in MPI mode). Defaults to zero. A good warm start (e.g. the previous solution in a time-stepping or optimization loop) cuts iterations; it does not change the converged solution or its gradients. Note that with the default RELATIVE_INI convergence the tolerance is relative to the initial residual, so a very good x0 tightens the target; consider convergence="ABSOLUTE" for warm-started loops.

None
config dict | None

AmgX configuration dictionary (see Solver Configuration for details). If None, JAX-AMG defaults are used.

None
block_dim int

Treat the matrix as a block matrix with square block_dim x block_dim blocks (e.g. coupled multi-component PDE systems with node-major interleaved unknowns: row i*block_dim + c is component c of node i). A and b keep their ordinary scalar CSR/vector form; the conversion to AmgX's BSR format happens internally. Rows must be divisible by block_dim (each rank's local partition in MPI mode). Since AmgX's classical AMG does not support blocks, the AMG defaults switch to aggregation (SIZE_2 + BLOCK_JACOBI); explicitly configured CLASSICAL AMG is rejected. In MPI mode the aggregation defaults use block-Jacobi coarse sweeps instead of DENSE_LU_SOLVER (which is broken in AmgX for distributed block matrices on 3+ ranks and rejected if configured explicitly); see Solver Configuration.

1
comm Comm | None

MPI communicator (typically mpi4py.MPI.COMM_WORLD). If provided, the solve runs in MPI mode. If not provided, MPI mode can still be used if MPI metadata has already been attached via with_cache(..., mpi=...).

None
nglobal int | None

Global matrix row count for MPI mode. Required when comm is provided and MPI metadata is not pre-attached to A.

None
partition_info tuple[int, int] | None

(row_start, row_end) owned by this rank in MPI mode. Required when comm is provided and MPI metadata is not pre-attached to A.

None
save_stats_file str | PathLike | None

Optional file path to save detailed AmgX solver statistics. If None, no file is created.

None
reuse_setup bool

For repeated solves with the same sparsity pattern, skip warm AMGX_solver_resetup and keep the cached hierarchy. This is cheaper per solve but may require more iterations if matrix coefficients change significantly.

False
**kwargs Any

Additional AmgX config parameters. These override values in config when both are provided.

{}

Returns:

Name Type Description
x Array

Solution vector (float32 or float64). In MPI mode, returns local portion.

info dict

Dictionary containing iterations, residual, status, and residual_history (residual norm per outer iteration, entry 0 being the initial residual; inside jit it has fixed length max_iters + 1 with NaN padding past entry iterations).

Status Codes

jaxamg.AMGXStatus

Bases: IntEnum

High-level AmgX solve status codes returned in info["status"] after calling jaxamg.solve.

These values are mapped from the native backend status for quick checks in Python code and in docs.

Members
  • SUCCESS: Solve converged successfully.
  • FAILED: Solver failed due to an internal/runtime error.
  • DIVERGED: Iterations diverged.
  • NOT_CONVERGED: Reached stopping criteria without convergence.

Caching

jaxamg.with_cache(A, *, coloring=None, mpi=None, is_symmetric=False)

Attach cached metadata (coloring, MPI info, or symmetry) to a matrix or operator.

This cache allows using matrices/operators inside JIT-compiled functions without recomputing metadata or passing it as separate arguments. See Caching Guide for more details.

Parameters:

Name Type Description Default
A MatrixOrOperator

A matrix or operator.

required
coloring tuple[ndarray, ndarray, ndarray, int, tuple[int, int]] | None

Cached coloring information from cache_coloring().

None
mpi dict[str, Any] | None

Cached MPI metadata from cache_mpi_metadata().

None
is_symmetric bool

If True, indicates the matrix is symmetric, allowing optimizations like skipping transpose in backward pass.

False

Returns:

Type Description
MatrixOrOperator

The same matrix/operator with requested cache attached.

jaxamg.cache_coloring(operator, shape)

Compute and cache coloring information for a callable operator.

Detection uses two methods, so the result is correct for ANY operator:

  1. Tracing: interpret the operator's jaxpr to recover the EXACT sparsity in a single trace (no probing), then colour and materialise it. Works for any JAX-expressed operator; skipped for operators that can't be traced structurally (opaque calls, data-dependent indexing).
  2. Probing (probe_sparsity_pattern + get_column_coloring): exhaustive one-hot basis-vector probing, correct for any operator -- the fallback when tracing is unavailable.

Parameters:

Name Type Description Default
operator Any

A callable operator A(x) that returns A @ x.

required
shape tuple[int, int] | int

Shape of the operator (n, m) or int size (for an n×n matrix). For a distributed operator this is the local block (n_local, n_global).

required

Returns:

Type Description
tuple[ndarray, ndarray, ndarray, int, tuple[int, int]]

Cached coloring information for reattachment with with_cache(..., coloring=...).

jaxamg.cache_mpi_metadata(config, comm, nglobal, partition_info, A, is_symmetric=False, save_stats=False, block_dim=1)

Pre-compute and cache MPI metadata for JIT-compatible solver usage.

The cached metadata can be reused across multiple JIT-compiled function calls with different matrices or operators (same structure).

Note

This function performs all non-traceable MPI operations outside the JIT boundary:

  • Computes static MPI communication metadata (recvcounts, displs)
  • Prepares MPI communicator pointer and local rank
  • Prepares config string
  • Computes max nnz across all ranks

Parameters:

Name Type Description Default
config dict

AmgX configuration dict or string

required
comm Comm

MPI communicator (from mpi4py.MPI.COMM_WORLD)

required
nglobal int

Global matrix size (total rows across all ranks)

required
partition_info tuple[int, int]

tuple (row_start, row_end) indicating which rows this rank owns

required
A MatrixOrOperator

Matrix or operator to compute max nnz for buffer sizing

required
is_symmetric bool

If True, the backward pass never transposes, so the transpose output size (nnz_out) is left unset (None). Should match the is_symmetric passed to with_cache; the default (False) computes it, which is always safe.

False
save_stats bool

If True, prepare the config with solver statistics output enabled, so a later solve(..., save_stats_file=...) on the cached matrix produces a complete stats file.

False
block_dim int

BSR block size for AmgX (see jaxamg.solve). Each rank's local partition must be divisible by it.

1

Returns:

Type Description
dict[str, Any]

A dictionary containing MPI metadata.

Note

The returned dictionary includes the following keys:

  • recvcounts_tuple: Tuple of row counts per rank
  • comm_ptr: MPI communicator pointer
  • lrank: Local GPU rank
  • nglobal: Global matrix size
  • config_str: Prepared configuration string
  • max_nnz: Maximum nnz across all ranks
  • nnz_out: This rank's local nnz(A^T) for the transpose output, or None when is_symmetric is True
  • halo_plan: Backward-pass halo-exchange plan for the gradient w.r.t. A (fetches only referenced remote solution entries)

Preconditioner

jaxamg.make_preconditioner(A, config=None, *, comm=None, nglobal=None, partition_info=None, save_stats_file=None, return_info=False, **kwargs)

Create a callable approximate inverse for external Krylov solvers.

The returned callable can be passed directly as the M argument to jax.scipy.sparse.linalg.cg(...) or jax.scipy.sparse.linalg.bicgstab(...).

By default the approximate inverse is a single AMG V-cycle (solver="AMG", max_iters=1), so each application is one cheap AMG sweep. This is deliberately different from jaxamg.solve, whose default is a full Krylov solve (PBICGSTAB) preconditioned by AMG: here AMG is the preconditioner and the outer Krylov method owns the iteration. Pass config/kwargs for a stronger inner application (e.g. more sweeps, a W-cycle, or max_iters=2).

Parameters:

Name Type Description Default
A MatrixOrOperator

Matrix or callable operator to precondition.

required
config dict[str, Any] | None

Optional AmgX configuration. If omitted, a single-cycle AMG approximate-inverse config is used.

None
comm Comm | None

Optional MPI communicator for distributed solves. If A already has MPI metadata attached via jaxamg.with_cache(..., mpi=...), this may be omitted.

None
nglobal int | None

Global matrix row count for MPI mode.

None
partition_info tuple[int, int] | None

Local row partition (row_start, row_end) for MPI mode.

None
save_stats_file str | None

Optional stats output path passed to jaxamg.solve(...).

None
return_info bool

If True, the returned callable yields (x, info) instead of only x.

False
**kwargs Any

Additional solver config overrides.

{}

Returns:

Type Description
Callable

A callable representing an approximate inverse M^{-1}.

jaxamg.make_lineax_preconditioner(operator, config=None, *, tags=_INHERIT_TAGS, comm=None, nglobal=None, partition_info=None, save_stats_file=None, **kwargs)

Wrap a Lineax operator as an AMG preconditioner operator.

This is the operator->operator counterpart of make_preconditioner: it maps a system operator A (a lineax.AbstractLinearOperator) to a preconditioner operator M with M.mv(r) ≈ A⁻¹ r, ready to hand to a Lineax solver via options={"preconditioner": M}. It folds the usual make_preconditioner plus FunctionLinearOperator wiring into a single call.

The operator's matrix-free action (operator.mv) is handed to JAX-AMG, whose sparsity detection assembles the explicit matrix AmgX needs (traced in one pass when possible, probed otherwise). The pattern is detected and cached eagerly here, since Lineax solvers apply the preconditioner under jax.jit where on-the-fly detection is impossible. A MatrixLinearOperator is assembled directly from its concrete matrix instead.

Parameters:

Name Type Description Default
operator AbstractLinearOperator

The system operator to precondition.

required
config dict[str, Any] | None

Optional AmgX configuration (see make_preconditioner).

None
tags Any

Lineax tags for the returned preconditioner. By default the operator's own tags are inherited (A⁻¹ shares A's symmetry/definiteness), which CG needs in order to accept the preconditioner. Pass an explicit value (e.g. ()) to override.

_INHERIT_TAGS
comm Comm | None

Optional MPI communicator for distributed solves.

None
nglobal int | None

Global matrix row count for MPI mode.

None
partition_info tuple[int, int] | None

Local row partition (row_start, row_end) for MPI mode.

None
save_stats_file str | None

Optional stats output path passed to jaxamg.solve(...).

None
**kwargs Any

Additional solver config overrides forwarded to make_preconditioner.

{}

Returns:

Type Description
FunctionLinearOperator

A lineax.FunctionLinearOperator approximating A⁻¹.

Runtime Utilities

jaxamg.get_solver_cache_info()

Inspect the internal C++ AmgX solver caches.

Returns:

Type Description
dict[str, Any]

A dictionary with cache size/capacity and entry summaries

dict[str, Any]

for single-GPU and MPI caches, plus whether isolated mode

dict[str, Any]

(JAXAMG_CACHE_SIZE=0) is active.

jaxamg.clear_solver_cache()

Clear the internal C++ AmgX solver cache. This releases all cached AmgX resources (matrices, solvers, vectors).

jaxamg.finalize()

Manually finalize AmgX resources. This clears the cache and calls AMGX_finalize. Normally only needed to be called manually in MPI mode to avoid shutdown-time resource warnings.