Getting started¶
Installation¶
Install the JAX-native core from PyPI:
pip install solvax
The optional native bridge needs SciPy:
pip install "solvax[native]"
Install a JAX accelerator build separately, following the JAX instructions for the target CUDA or TPU platform. SOLVAX does not select the JAX backend.
The v0.7.3 wheel built from this tree is 54 KiB. Core dependencies are JAX and Equinox; SciPy remains optional for the host-side native bridge. SOLVAX does not depend on a second linear-operator package, which keeps installation and operator ownership explicit.
The operator contract¶
Most iterative SOLVAX routines accept a callable matvec implementing
The matrix does not need to exist as an array:
def matvec(v):
return diffusion(v) + advection(v) + reaction * v
This contract supports dense arrays, stencil applications, spectral
transforms, and nested solver calls without changing the outer method. GMRES
and PCG accept arbitrary JAX pytrees, provided matvec, precond, b, and
x0 all preserve the same tree structure. GCROT currently operates on flat
one-dimensional arrays.
Preconditioner contract¶
The precond argument is an inverse action:
not the matrix \(M\). SOLVAX FGMRES uses right preconditioning,
while PCG applies a positive-definite preconditioner to the residual. A direct SOLVAX factorization can therefore become an iterative preconditioner by closing over its factors:
factors = sx.block_thomas_factor(lower, diagonal, upper)
def precondition(r):
shaped = r.reshape(n_blocks, block_size)
return sx.block_thomas_solve(factors, shaped).reshape(-1)
See Preconditioners for construction patterns.
Tolerances and stopping¶
The linear iterative solvers use the residual criterion
Consequences:
Use
atol > 0when a zero or tiny right-hand side is meaningful.A small relative tolerance cannot overcome an inaccurate or unstable preconditioner.
Always inspect
convergedand the true final residual; reaching an iteration limit is a valid result state, not an exception.
Fixed-point acceleration uses the true map residual \(\lVert G(x)-x\rVert_2\) and scales its relative tolerance with \(\max(\lVert x_0\rVert_2,1)\).
Shape conventions¶
Method |
Principal shape convention |
|---|---|
|
|
|
system dimension first; all trailing axes are batched |
|
real |
|
array or arbitrary matching pytree |
|
flat |
|
history on axis 0 |
|
same output layout as the corresponding JAX transform |
For block tridiagonal storage, lower[0] and upper[-1] are unused but must be
present. For banded storage, the main diagonal is row upper_bw, matching the
SciPy banded convention.
Result objects¶
Krylov results¶
gmres and gcrot return KrylovSolution:
solution.x
solution.residual_norm
solution.iterations
solution.converged
solution.recycle # None for GMRES; (C, U) for GCROT
PCG results¶
pcg and pcg_linear_solve return PCGSolution, adding relative residual,
status, and a fixed-shape residual history. Convert a materialized status to a
name with status_name outside a JAX trace.
Fixed-point results¶
aitken_fixed_point returns the iterate, true residual norm, iteration count,
convergence flag, and final relaxation parameter.
Newton–Krylov results¶
newton_krylov returns NewtonKrylovSolution, containing the final PyTree
iterate, true nonlinear residual norm, Newton and total GMRES iteration counts,
and separate nonlinear and linear convergence flags. The solve is matrix-free;
pass precond= for a right preconditioner and inner_product= for weighted or
distributed Krylov products.
JAX transforms¶
All core routines are intended for jit; static algorithm sizes such as
restart, m, k, max_steps, and refinement counts determine compiled
shapes and should not vary inside a trace.
solve = jax.jit(lambda rhs: sx.gmres(matvec, rhs, restart=30).x)
batched = jax.vmap(solve)(stacked_rhs)
Differentiating directly through a converged iterative method also
differentiates its algorithm. For gradients of the mathematical solution, use
linear_solve, root_solve, or pcg_linear_solve; see
Matrix-free nonlinear solves and implicit differentiation.
Checking a solve¶
During development, check the residual independently:
x = solution.x
absolute = jnp.linalg.norm(b - matvec(x))
relative = absolute / jnp.maximum(jnp.linalg.norm(b), jnp.finfo(b.dtype).tiny)
For a structured direct solver, also compare a small instance against
jnp.linalg.solve. A small dense reference is a test oracle, not a production
strategy.