Skip to main content

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

Project description

SOLVAX

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, p-multigrid, Kronecker approximations, symmetric additive and line smoothers), mixed-precision iterative refinement, and implicit differentiation of every solve — all jit/vmap/grad-transparent, on CPU and GPU.

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.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.
generated_factors = sx.block_thomas_factor_fn(block_fn, n_blocks=N)

# 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)

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, p-multigrid V-cycles, nearest-Kronecker, mixed-precision wrappers
solvax.direct Block-tridiagonal Schur elimination (block Thomas): full, factor/solve split, truncated-storage mode
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-style Krylov subspace recycling for parameter continuation
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: harmonic-Ritz recycle selection, pytree GCROT operands, 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) tridiagonal line, corners in lower[0] and upper[-1]:
x = sx.cyclic_tridiagonal_solve(lower, diag, upper, 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.

Project details


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.8.7.tar.gz (297.7 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.8.7-py3-none-any.whl (72.8 kB view details)

Uploaded Python 3

File details

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

File metadata

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

File hashes

Hashes for solvax-0.8.7.tar.gz
Algorithm Hash digest
SHA256 53af93fbc227898c776ddb0505dd537d4273449d1c598f32bb7a04bdca6c93ee
MD5 c3d0f37618198e2f04a644a346901a7e
BLAKE2b-256 205f5a9e276715397f92a3fa9f420c4dba563ce0a5f165b42d4b836b4ba51411

See more details on using hashes here.

Provenance

The following attestation bundles were made for solvax-0.8.7.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.8.7-py3-none-any.whl.

File metadata

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

File hashes

Hashes for solvax-0.8.7-py3-none-any.whl
Algorithm Hash digest
SHA256 1d15b0d6798ae12ed32a7c76bee1cee1ea3bc739a700f2b1a86b8baa1eeefa42
MD5 e869fb9ae37dcc8654608d4ac495ffda
BLAKE2b-256 f11a3fd574bd6bdef0aff3ccfca328e3a28d2e380b512a76fd919fa2af1b1999

See more details on using hashes here.

Provenance

The following attestation bundles were made for solvax-0.8.7-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.

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