nlls_gram
Metric-aware Levenberg-Marquardt nonlinear least-squares for JAX pytrees, aimed at solving systems of equations (i.e., interpolation). The core use cases are underdetermined systems — more parameters than residuals — where LM is solved in its Gram form, and square or redundant tall nonlinear systems where regularization is needed to select among the interpolating solutions, with the selection controlled by a user-chosen metric.
LevenbergMarquardt minimizes a user residual taking
(x), (x, args), or (x, args, p), always in that order. The unknown x
may be any JAX pytree; internally it
is flattened with jax.flatten_util.ravel_pytree. The default solver picks
the smaller dense factorization from the problem shape — the residual-space
Gram system or the whitened normal system — with QR, CG, and matrix-free
LSMR alternatives. Use update(...) for a single LM step or solve(...)
for an internally jitted loop.
Problem
At each iteration the solver builds a step (s) from the metric-damped linearized subproblem
$$ \min_s \frac12|r + Js|_2^2 + \frac{\lambda}{2}s^\top M s, \qquad M \succ 0. $$
The default is (M=I). For kernel/RKHS coefficient problems, if
$$ f_\alpha(x)=\sum_{j=1}^n \alpha_j K(x,x_j), $$
then
$$ |f_\alpha|_{\mathcal H_K}^2 = \alpha^\top K\alpha, $$
so the natural parameter metric is (M=K), not the Euclidean metric.
Near the interpolation threshold the residuals are small, damping falls, and LM becomes metric Gauss-Newton: each step is the minimum-(M)-norm correction solving the linearized residual equations,
$$ s = -M^{-1}J^\top\left(JM^{-1}J^\top\right)^{-1}r = \arg\min_s |s|_M ;;\text{s.t.};; r + Js = 0. $$
With an RKHS metric this selects minimum-RKHS-norm corrections — kernel methods let you control exactly which norm that is. The docs derive this and the large-damping limit (metric gradient descent).
Install
uv add nlls-gram
For GPU use, install the JAX accelerator build that matches your hardware, for example:
uv add nlls-gram "jax[cuda13]"
Minimal Example
import jax
import jax.numpy as jnp
from nlls_gram import LevenbergMarquardt
def residual_fn(x, args):
ts, ys = args
return x["a"] * jnp.exp(x["b"] * ts) - ys
ts = jnp.linspace(0.0, 2.0, 20)
ys = 2.0 * jnp.exp(-1.0 * ts)
x = {"a": 1.0, "b": 0.0}
solver = LevenbergMarquardt(residual_fn, init_damping=1e-2)
lm_state = solver.init(x, (ts, ys))
@jax.jit
def train_step(x, lm_state):
return solver.update(x, lm_state, (ts, ys))
for _ in range(50):
x, lm_state, info = train_step(x, lm_state)
print(x["a"], x["b"]) # approximately 2.0, -1.0
For a simple full solve loop:
result = solver.solve(x, (ts, ys), max_steps=50, atol=1e-8)
x = result.x
solve stops on a residual-norm atol, gradient-norm gtol, or
accepted-step-norm xtol (each 0.0 disables), always enforces max_steps,
and takes a traceable callback for custom stopping, epoch-style data
resampling, and per-step history recording; the docs have a cookbook.
solve(...).x also supports custom implicit JVP/VJP with respect to p;
the docs give the metric-minimum-norm formula and a minimal jax.jvp /
jax.vjp example. The default ad_solver="auto" uses a direct solve for
every square system, preserves the forward CG space for nonsquare CG systems,
and uses SVD otherwise. Every method is independently swappable (an lsmr forward solve with
ad_solver="normal_cg" is fully matrix-free end to end). The metric
matters for underdetermined roots because it selects which tangent is the
minimum-norm solution. The per-step update(...) interface does not define
the implicit AD rule. By default, both CONVERGED and MAX_STEPS results are
usable: fixed-step solves retain their implicit derivative while the status
still reports MAX_STEPS. Pass max_steps_is_success=False for strict
failure semantics. A failed solve keeps its primal result and diagnostics but
contributes exactly zero through result.x and result.aux; result.p
remains an identity pass-through. Its linear tangent program is evaluated
safely at differentiation-inert copies of the caller's original
(x0, args, p), so those initial inputs must be JVP-safe for the residual,
aux map, and any metric or preconditioner factory used by the selected AD
method. This also keeps mixed successful/failed vmap lanes finite.
Metric Example
For a dense SPD metric (M = LL^\top), use the Cholesky helper:
import jax.numpy as jnp
from nlls_gram import LevenbergMarquardt, metric_from_cholesky
L = jnp.linalg.cholesky(metric_matrix)
solver = LevenbergMarquardt(
residual_fn,
init_damping=1e-2,
metric=metric_from_cholesky(L),
)
The Metric callbacks act on the flattened parameter vector. The docs give
the exact callback contract, branch formulas, and validation rules, plus
structural constructors (metric_from_tridiagonal_precision,
metric_from_state_space and metric_from_quasiseparable for exact O(n)
Matérn/state-space kernel Grams, metric_from_diagonal,
blockdiag_metric) so common metrics need no callback plumbing.
For a metric that depends on the current iterate or on residual aux outputs
(has_aux=True), pass metric_factory=MetricFactory(prepare, build) instead
of metric: prepare(x, args, p, aux) caches a state once per accepted step
and build(state) returns a plain Metric — any of the constructors above
slots in directly, e.g.
MetricFactory(prepare=lambda x, args, p, aux: aux["L"], build=metric_from_cholesky).
Solvers
linear_solver="auto"(the default): resolves at trace time to the smaller dense factorization —gram_choleskywhenn > m,normal_choleskyotherwise. A shape rule, and safely so: the two forms compute the same step.linear_solver="gram_cholesky": densem × mresidual-space Gram solve.linear_solver="normal_cholesky": densen × nwhitened normal solve; its small-damping limit is the minimum-metric-norm least-squares step at every shape and rank.linear_solver="qr": dense QR solve of the whitened-step problem (requires a full-row-rank Jacobian).linear_solver="augmented_qr": direct augmented QR in parameter space; robust to rank-deficient Jacobians when damping is positive and best suited to small systems.linear_solver="gram_cg": matrix-free residual-space CG. Adual_preconditioneris required (e.g.sherman_morrison_preconditioner, or the randomizednystrom_preconditionerfor neural-network duals; passidentity_preconditioner()to run unpreconditioned CG explicitly); on a nonsquare systemad_solver="auto"keepssolve(...).xmatrix-free under AD and requiresad_solver_preconditionerwhen the differentiated solve is traced (an explicitad_solver="gram_cg"validates it eagerly). When the dual operator rotates as LM driftsx, passpreconditioner_factory=PreconditionerFactory(prepare, apply)instead — a θ-adaptive preconditioner rebuilt from the live iterate each step — andrecycle=RecycleConfig(rank=k)to carry a deflation basis across steps, recycling each solve's Krylov subspace into the next.linear_solver="normal_cg": matrix-free CG on the whitened normal system, iterating in parameter space — the matrix-free form for square-to-tall problems. Anormal_preconditioneris required; on rank-deficient problems it must preserverange(Bᵀ)or the minimum-norm selection is lost (identity_preconditioner()always qualifies — the docs give the full requirement).linear_solver="lsmr": matrix-free LSMR on the whitened augmented system, the iterative sibling ofaugmented_qr, using only J/Jᵀ products. It works on the whitened Jacobian rather than a squared Gram/normal operator, so it stays accurate at small damping where those solves hit theireps·condfloor. An optionalwhitened_preconditioner=WhitenedPreconditioner(solve, solve_transpose)right-preconditions the operator to cluster its spectrum; every damped posed subproblem stays exactly the identity-damped whitened one, so the preconditioner changes the iteration path, never the converged step.
All eight solve the same metric-damped linearized subproblem up to the
accuracy of the chosen linear solver. The dense paths materialize the
Jacobian from its small side (jacobian_mode="auto"; "fwd"/"rev"
force one AD mode), so tall systems never build an m × m residual basis.
Docs and Alternatives
Full docs: https://highdimensionaleconlab.github.io/nlls_gram/
Working with an AI assistant? Point it at
docs/tuning_guide.md
if it doesn't pick it up automatically — solver selection, damping heuristics,
inner-solve scheduling, and failure signatures, written to be read by humans
and agents alike (also indexed via the site's llms.txt).
For a broader JAX nonlinear solver library, see
Optimistix. nlls_gram is more
specialized: it focuses on underdetermined nonlinear least-squares, residual
space Gram solves, and explicit parameter-space metrics.
License
MIT
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file nlls_gram-2.4.0.tar.gz.
File metadata
- Download URL: nlls_gram-2.4.0.tar.gz
- Upload date:
- Size: 2.0 MB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d6b7e42b4dc3d05b7a388ccb37b7b09baa5982d10aead72ea010572bf93c4248
|
|
| MD5 |
3368d62520b03bb00bb9fb806ab997cd
|
|
| BLAKE2b-256 |
bccd3195db6a86d1e40b703a12fb796924aa2eeb703ad65eb1ee22585e431edf
|
Provenance
The following attestation bundles were made for nlls_gram-2.4.0.tar.gz:
Publisher:
publish.yml on HighDimensionalEconLab/nlls_gram
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
nlls_gram-2.4.0.tar.gz -
Subject digest:
d6b7e42b4dc3d05b7a388ccb37b7b09baa5982d10aead72ea010572bf93c4248 - Sigstore transparency entry: 2214941504
- Sigstore integration time:
-
Permalink:
HighDimensionalEconLab/nlls_gram@e1f163e91594d5a625999396b995b184885e8707 -
Branch / Tag:
refs/tags/v2.4.0 - Owner: https://github.com/HighDimensionalEconLab
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@e1f163e91594d5a625999396b995b184885e8707 -
Trigger Event:
release
-
Statement type:
File details
Details for the file nlls_gram-2.4.0-py3-none-any.whl.
File metadata
- Download URL: nlls_gram-2.4.0-py3-none-any.whl
- Upload date:
- Size: 77.0 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via:
twine/6.1.0 CPython/3.13.14
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
d2beceeb15d6e0668d4ad63e4a04c69b7a6a5f88f8863171c77454743930f7c7
|
|
| MD5 |
3f49cf303f66a204b787d5a6f27f63b4
|
|
| BLAKE2b-256 |
c37ddd2e404450bdb687669a35ae45ecb15f04e0932bdde547e31722b2d60b53
|
Provenance
The following attestation bundles were made for nlls_gram-2.4.0-py3-none-any.whl:
Publisher:
publish.yml on HighDimensionalEconLab/nlls_gram
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
nlls_gram-2.4.0-py3-none-any.whl -
Subject digest:
d2beceeb15d6e0668d4ad63e4a04c69b7a6a5f88f8863171c77454743930f7c7 - Sigstore transparency entry: 2214941514
- Sigstore integration time:
-
Permalink:
HighDimensionalEconLab/nlls_gram@e1f163e91594d5a625999396b995b184885e8707 -
Branch / Tag:
refs/tags/v2.4.0 - Owner: https://github.com/HighDimensionalEconLab
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
publish.yml@e1f163e91594d5a625999396b995b184885e8707 -
Trigger Event:
release
-
Statement type: