Skip to main content

SOLVAX

DOI

tests codecov PyPI docs license

Differentiable structured linear solvers, preconditioners and matrix-free methods in JAX.

solvax provides the solver infrastructure that kinetic and PDE codes keep re-implementing: structured direct solves (batched dense LU, block-tridiagonal Schur elimination with truncated storage), preconditioned and recycled Krylov methods, physics-agnostic preconditioners (coarse-operator LU, semicoarsened geometric and p-multigrid, Kronecker approximations, symmetric additive and line smoothers), mixed-precision iterative refinement, and implicit differentiation of every solve — traceable under jit, vmap and grad, on CPU and GPU.

Two documented exceptions, because "transparent to every transform" would not be true: the exact-window reverse rule is a custom_vjp, so jax.jacfwd and jax.jvp raise on it rather than falling back to the taped path (pass the full window, or differentiate the untruncated entry point, when forward mode is what you need); and solvax.native runs SciPy's SuperLU on the host, so it refuses to be traced at all and says so.

It complements general JAX solver libraries with block-structured direct elimination, coarse-operator and multigrid preconditioning, and Krylov subspace recycling for parameter continuation. SOLVAX operators are native JAX pytrees; no external operator abstraction is required.

Install

pip install solvax

Quickstart

import jax
import jax.numpy as jnp
import solvax as sx

# Solve a block-tridiagonal system L_k x_{k-1} + D_k x_k + U_k x_{k+1} = b_k
x = sx.block_thomas(lower, diag, upper, rhs)

# Matrix-free PCG on arrays or arbitrary JAX pytrees
solution = sx.pcg(matvec, rhs, precond=preconditioner, rtol=1e-10)
assert solution.converged

# Solve an expensive affine coupling map without assembling its Jacobian
coupled = sx.affine_fixed_point_gmres(coupling_sweep, initial_state)

# Same diagnostics, but gradients use an implicit primal/transpose solve
implicit_solution = sx.pcg_linear_solve(matvec, rhs, precond=preconditioner)

# Reuse one elimination across many right-hand sides
factors = sx.block_thomas_factor(lower, diag, upper)
x1 = sx.block_thomas_solve(factors, rhs1)
x2 = sx.block_thomas_solve(factors, rhs2)

# Generate each block once when reusable factors are needed without a stored
# diagonal band. `row(j)` returns the triple (L_j, D_j, U_j) of row j; the
# parameterized form `block_fn(params, j)` used further down takes the
# parameters as its first argument.
generated_factors = sx.block_thomas_factor_fn(row, n_blocks=N)

# One generated solve with O(sqrt(N) m^2) factor storage and exact JVP/VJP.
x = sx.block_thomas_checkpointed_fn(row, N, rhs)

# Memory-truncated mode: rhs nonzero only in the lowest K blocks and only the
# lowest K solution blocks needed -> O(K m^2) memory, independent of N.
x_low = sx.block_thomas_truncated(lower, diag, upper, rhs[:3], keep_lowest=3)

Differentiate a generated selected-head solve with respect to the compact parameters that build its rows, at retained state independent of the block count. The window is estimated up front from the chain's own localization profile. The estimate is a diagnostic, not a certificate: it reports where the chain's transfer norms drop below one, and you should confirm the accuracy you need by widening the window until the gradient stops moving.

advice = sx.localization_crossover_window(lambda k: block_fn(p, k), N, keep_lowest=3)
# advice.certified is False: an estimate, not a guarantee. It can be passed
# straight back to the solver, or unpacked as advice.window.

def objective(params):
    x_low = sx.block_thomas_truncated_fn(
        block_fn, N, rhs[:3], keep_lowest=3,
        params=params, adjoint_window=advice,
    )
    return loss(x_low)

grad = jax.grad(objective)(p)

# Confirm the window before trusting it: widen it and see if the gradient moves.
report = sx.check_localized_gradient(
    lambda w: jax.grad(lambda q: loss(sx.block_thomas_truncated_fn(
        block_fn, N, rhs[:3], keep_lowest=3, params=q, adjoint_window=w)))(p),
    window=advice.window,
)

Everything is differentiable (jax.grad through the solve) and batchable (jax.vmap over stacked systems).

What's in the box

Module Contents
solvax.operators Matrix-free, sum, Kronecker, block-tridiagonal and bordered (constraint-row) operator containers with closed-form transposes
solvax.precond Jacobi/block-Jacobi, coarse-operator LU, Galerkin-deflation coarse correction, symmetric additive and alternating-direction line composition, V-/W-/F-cycle multigrid over explicit or semicoarsened rediscretized hierarchies, nearest-Kronecker, mixed-precision wrappers
solvax.transfer Separable per-axis restriction/prolongation (full weighting, linear, injection) with periodic, dirichlet and reflective closures, exact variational adjointness, and semicoarsening plans
solvax.smoothers Point/block Jacobi, batched tridiagonal line and exact banded plane relaxation, upwind-ordered sweeps for streaming operators, and a measured smoothing factor
solvax.direct Block-tridiagonal Schur elimination (block Thomas): full, factor/solve split, selected-head (truncated-storage) mode, exact-window localized adjoint, per-row localization profile and window advisor
solvax.banded Non-pivoted banded LU with row equilibration + static pivoting; periodic variant via the Woodbury capacitance trick
solvax.tridiagonal Batched scalar tridiagonal solve (reproducible Thomas / fused cuSPARSE backend) and periodic (cyclic) systems via a Sherman--Morrison correction
solvax.elliptic Spectral Fourier--Helmholtz solve for separable periodic-by-bounded elliptic problems — the drift-plane / vorticity lap phi = rhs inversion, one FFT + one batched tridiagonal sweep
solvax.krylov Flexible restarted GMRES (CGS2 + Givens) over arrays, scalars and arbitrary pytrees with optional custom inner products, and GCROT Krylov subspace recycling with FIFO or harmonic-Ritz (GCRO-DR) deflated restarting
solvax.pcg Matrix-free pytree PCG with preconditioning, fixed-shape residual history, and explicit convergence/breakdown status
solvax.fixed_point Safeguarded Aitken, bounded-memory (condition-filtered) Anderson, and matrix-free affine fixed-point FGMRES
solvax.implicit Matrix-free newton_krylov (JFNK) plus implicit-function-theorem linear_solve and root_solve — gradients cost one extra (transposed) solve
solvax.autodiff Bounded-memory chunked forward/reverse Jacobians (chunked_jacfwd/jacrev/jacobian) with automatic chunk sizing
solvax.refine Mixed-precision iterative refinement (float32 factor, float64 residuals)
solvax.native Host-side SuperLU bridge (non-differentiable, import-guarded)

Complex-valued GMRES/GCROT, tridiagonal solves, and fixed-point acceleration use Hermitian inner products and real-valued safeguards. Remaining roadmap: multi-leaf pytree GCROT operands (GCROT takes arrays of any rank; GMRES is pytree-native) and expanded GPU batched-LU benchmarks.

# Preconditioned, recycled Krylov across a parameter scan:
sol = sx.gcrot(matvec, b, precond=coarse_inverse, m=50, k=10)
sol2 = sx.gcrot(matvec2, b2, precond=coarse_inverse, recycle=sol.recycle)

# Matrix-free Newton-Krylov (JFNK): Jacobian-vector products via jax.linearize,
# each correction solved by FGMRES over an array or structured pytree state.
root = sx.newton_krylov(residual_fn, x0, precond=approx_inverse, rtol=1e-8)

# Weakly contractive affine coupling map G(x) = L x + c, solved as (I - L) x = c:
fixed = sx.affine_fixed_point_gmres(coupling_map, x0, restart=20)

# Periodic (cyclic) scalar tridiagonal line, corners in sub[0] and sup[-1]:
x_line = sx.cyclic_tridiagonal_solve(sub, dia, sup, line_rhs)

# Differentiable solve wrapping any solver:
x = sx.linear_solve(matvec, b, solver=lambda mv, rhs: sx.gmres(mv, rhs).x)

License

MIT. Developed by the UW Plasma group.

Download files

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

Source Distribution

solvax-0.11.1.tar.gz (395.3 kB view details)

Uploaded Source

Built Distribution

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

solvax-0.11.1-py3-none-any.whl (111.5 kB view details)

Uploaded Python 3

File details

Details for the file solvax-0.11.1.tar.gz.

File metadata

  • Download URL: solvax-0.11.1.tar.gz
  • Upload date:
  • Size: 395.3 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.14

File hashes

Hashes for solvax-0.11.1.tar.gz
Algorithm Hash digest
SHA256 6031763040e7a43fe5a3466be3635f4ec61996d4574388b67a02a221cc87c1b8
MD5 437ad9ff0ad9f0644e6a7225113d3522
BLAKE2b-256 f6d51d68200c7f46e2ac48c6fd1924721f42b6c1f67beb75219ad9f8d9c8b678

See more details on using hashes here.

Provenance

The following attestation bundles were made for solvax-0.11.1.tar.gz:

Publisher: publish.yml on uwplasma/SOLVAX

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file solvax-0.11.1-py3-none-any.whl.

File metadata

  • Download URL: solvax-0.11.1-py3-none-any.whl
  • Upload date:
  • Size: 111.5 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.14

File hashes

Hashes for solvax-0.11.1-py3-none-any.whl
Algorithm Hash digest
SHA256 7a84a24859fed2b231c89609b62de3a813f6676c7662243b3a70ec2949abe988
MD5 3ee9e93a0cde43cb2822c38b3872ee6b
BLAKE2b-256 edf2e2678817cb13c702fda14e09588b81a310be13e279f27b1cd790639d8fb1

See more details on using hashes here.

Provenance

The following attestation bundles were made for solvax-0.11.1-py3-none-any.whl:

Publisher: publish.yml on uwplasma/SOLVAX

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

0.22.0

2 files

0.21.0

2 files

0.20.0

2 files

0.19.0

2 files

0.18.0

2 files

0.17.0

2 files

0.16.0

2 files

0.15.0

2 files

0.14.0

2 files

0.13.0

2 files

0.12.0

2 files

0.11.2

2 files

This release

0.11.1 This release

2 files

0.11.0

2 files

0.10.1

2 files

0.10.0

2 files

0.9.1

2 files

0.9.0

2 files

0.8.8

2 files

0.8.7

2 files

0.8.6

2 files

0.8.4

2 files

0.8.3

2 files

0.8.2

2 files

0.8.1

2 files

0.8.0

2 files

0.7.3

2 files

0.7.2

2 files

0.7.1

2 files

0.7.0

2 files

0.6.1

2 files

0.6.0

2 files

0.5.1

2 files

0.5.0

2 files

0.2.0

2 files

0.1.0

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page