Additive Symbolic Regression#
Additive symbolic regression fits a model as a sum of small symbolic expressions:
f(x) = c + eta_1 * g_1(x) + eta_2 * g_2(x) + ... + eta_K * g_K(x)
where each g_k(x) is a small, interpretable symbolic expression discovered by
the existing JAXSR machinery. This is analogous to gradient boosting, except
each weak learner is a symbolic expression rather than a decision tree.
The submodule lives in jaxsr.additive.
Three flavours of symbolic regression#
Approach |
What it does |
Status |
|---|---|---|
Single-expression ( |
Fits one sparse expression over a fixed basis library. |
Available |
Stagewise additive ( |
Repeatedly fits a small expression to the residual and adds it to the ensemble. Old terms are frozen. |
Available |
Backfitting additive ( |
Maintains a fixed set of terms and revises each one in place across sweeps (GAM-style). |
Available (squared error); Bayesian variant planned |
The key distinction between the two additive variants:
Stagewise: once a term is discovered it never changes; only its linear weight may be re-estimated.
Backfitting: terms are revised repeatedly, each conditioned on the current fit of all the others.
Scope: what this is (and isn’t) good for#
In one line: JAXSR is a linear method over a fixed feature space — it selects a sparse combination of basis functions you supply. It is not a free-composition equation discoverer.
That distinction decides whether it is the right tool:
Good fit: the right building blocks are on the menu (or you can add them with
BasisLibrary.add_custom), and you want an interpretable, robust, uncertainty-aware additive model. On targets that live in the library it is fast and accurate, and the additive layer adds robust/quantile losses and structural-uncertainty bootstrapping that genetic-programming tools don’t offer out of the box.Wrong fit: you want to discover an unknown compositional law such as
exp(x0*x1),x0 / (1 + x1**2), orsin(2*x0). These are not single basis functions, and the space of such compositions is infinite and continuously parameterized, so no fixed library enumerates them in advance. For that, reach for a genetic-programming or neural symbolic-regression tool (PySR, Operon, AI-Feynman), which search the space of expressions instead of selecting from a fixed dictionary — or try the experimentalRecursiveSymbolicRegressor, which grows compositions along the residual and partially lifts this ceiling.
The limit is one of discovery, not representation: the linear-in-basis model
fits any of those targets perfectly the moment the exact term is in the library
(e.g. library.add_custom("exp(x0*x1)", lambda X: jnp.exp(X[:, 0] * X[:, 1]))) —
it simply cannot figure out which composition it needs without being told.
JAXSR’s parametric bases can additionally fit a few constants inside a
pre-specified nonlinearity (e.g. sin(a*x0) with a optimized), but that still
requires you to name the functional form. The additive extensions in this guide
raise the statistical sophistication (boosting, robust/quantile losses,
backfitting, structural UQ); they do not change this expressiveness boundary.
Quick start#
import numpy as np
from jaxsr.additive import StagewiseSymbolicRegressor
rng = np.random.default_rng(0)
X = rng.uniform(-2, 2, size=(200, 2))
y = 2.0 * X[:, 0] + 0.5 * X[:, 1] ** 2 + 0.1 * rng.normal(size=200)
model = StagewiseSymbolicRegressor(
n_terms=5,
learning_rate=0.2,
max_complexity=6,
refit_coefficients=True,
)
model.fit(X, y)
print(model) # pretty structural summary
print(model.expressions_) # per-term expression strings
print(model.coefficients_) # per-term weights
print(model.intercept_) # additive intercept
y_pred = model.predict(X)
The print(model) output looks like:
StagewiseSymbolicRegressor(
intercept = 1.07
terms =
+ 1 * (y = 2*x0 - 1.07 + 0.5*x1^2)
...
)
The stagewise algorithm#
Initialise the intercept to
mean(y)and the prediction to that constant.Compute the residual
y - prediction.Fit a small symbolic expression
g_kto the residual (viajaxsr.fit_symbolic).Append
g_kto the ensemble.If
refit_coefficients=True, rebuild the design matrixPhi[:, j] = g_j(X)and re-solvey ~= intercept + Phi @ coefficientsby least squares. Otherwise, updateprediction += learning_rate * g_k(X).Record train (and optional validation) loss.
Repeat until
n_termsterms are added or early stopping triggers.
Key parameters#
Parameter |
Meaning |
|---|---|
|
Maximum number of boosting stages (terms). |
|
Shrinkage on each stage when |
|
Complexity budget per term (max basis terms). Keep small to favour many simple terms. |
|
Re-solve all linear weights by OLS after each stage. |
|
|
|
Hold out a validation split and stop when it stops improving. |
|
Early-stopping controls. |
|
Which basis functions each term may use. |
|
Complexity control within each term ( |
Coefficient refitting#
refit_coefficients=False: the weights are the learning-rate-scaled stagewise weights (coefficients_[k] == learning_rate).refit_coefficients=True: after each new term, the intercept and all per-term weights are re-solved by ordinary least squares over the discovered symbolic features. This decouples term discovery (nonlinear, greedy) from term weighting (linear, global) and typically improves accuracy.
The refit uses jnp.linalg.lstsq (SVD-based, minimum-norm), so the highly
correlated columns produced by later boosting stages do not cause instability.
Combined expression#
model.to_expression() returns a single simplified SymPy expression combining
all terms:
expr = model.to_expression() # requires sympy
Saving and loading#
Fitted models serialize to JSON (each term is stored via the underlying
SymbolicRegressor state), mirroring the rest of jaxsr. Note that the models
are not picklable — the basis-function closures cannot be pickled — so use
save/load rather than pickle:
model.save("additive_model.json")
loaded = StagewiseSymbolicRegressor.load("additive_model.json")
Structural uncertainty (bootstrap)#
A single fitted expression can hide the fact that the structure itself is
uncertain — several different basis sets may explain the data about equally
well (this is common with collinear features). bootstrap_additive refits the
model on bootstrap resamples and reports, for each basis function, how often it
is selected — a cheap approximation to a posterior inclusion probability —
together with a predictive ensemble:
from jaxsr.additive import (
StagewiseSymbolicRegressor,
bootstrap_additive,
bootstrap_predict_additive,
)
est = StagewiseSymbolicRegressor(n_terms=3, max_complexity=2)
res = bootstrap_additive(est, X, y, n_bootstrap=100, random_state=0)
# How stable is the discovered structure?
for name, prob in res["inclusion_probabilities"].items():
print(f"{name:10s} selected in {prob:.0%} of resamples")
# Prediction intervals that reflect *structural* variability, not just noise
pi = bootstrap_predict_additive(res["models"], X_new, alpha=0.1)
pi["mean"], pi["lower"], pi["upper"]
How to read it. Inclusion probabilities near 0 or 1 mean the structure is identifiable and the single fitted expression is trustworthy. Diffuse values (e.g. a basis selected 50–60% of the time) mean the data do not determine one expression — no single symbolic model should be over-trusted, and the bootstrap intervals are the honest summary. This also works as a decision gate for heavier Bayesian modelling: if the probabilities are already crisp, there is little structural uncertainty left to quantify. It works for both the stagewise and backfitting regressors.
Early stopping#
With early_stopping=True, a validation split (validation_fraction) is held
out. After each stage the validation loss is recorded; training stops once it
fails to improve by at least min_delta for patience consecutive stages, and
the model rolls back to the best iteration.
Losses: robust and quantile regression#
This is where additive symbolic regression goes beyond ordinary least-squares
symbolic regression. Each weak learner fits the negative gradient -dL/dy_pred
(gradient boosting), so you can target losses that OLS selection cannot:
|
Class |
Use when |
|---|---|---|
|
|
Standard regression |
|
|
Outliers present (fits the median) |
|
|
Outliers, but keep efficiency near zero |
|
|
Quantiles / prediction intervals / asymmetric cost |
Pass a name for defaults, or an instance to customise:
from jaxsr.additive import StagewiseSymbolicRegressor, QuantileLoss, HuberLoss
# Robust regression: heavy outliers barely move the fit
robust = StagewiseSymbolicRegressor(loss="huber", learning_rate=0.5).fit(X, y)
# 90th-percentile regression (build intervals by fitting several quantiles)
q90 = StagewiseSymbolicRegressor(loss=QuantileLoss(0.9), learning_rate=0.5).fit(X, y)
How non-squared losses are fit. Each stage fits a symbolic term to the
negative gradient, then a line search picks the step size that minimises the
loss (learning_rate shrinks that step). Because the ordinary least-squares
coefficient refit targets squared error, refit_coefficients=True is ignored for
non-squared losses (a warning is issued) and gradient boosting is used instead —
so set refit_coefficients=False explicitly for robust/quantile models.
The optimal constant initialisation adapts to the loss: mean for squared error, median for absolute/Huber, and the empirical quantile for quantile loss.
Add further losses (Poisson, logistic, …) by subclassing Loss and
registering them in jaxsr.additive.losses._LOSSES.
Backfitting (GAM-style)#
BackfittingSymbolicRegressor maintains a fixed number of terms and
revises each one across sweeps, instead of freezing them. Each sweep removes a
term, re-discovers its expression on the partial residual, and puts it back:
from jaxsr.additive import BackfittingSymbolicRegressor
model = BackfittingSymbolicRegressor(n_terms=4, n_sweeps=6, max_complexity=3)
model.fit(X, y) # warm-started from a stagewise fit, then refined by sweeps
for sweep in 1..n_sweeps:
for term j:
partial_residual = y - intercept - sum_{i != j} coef_i * g_i(X)
g_j = fit_symbolic(X, partial_residual, ...) # re-discover structure
intercept, coef = OLS refit over all terms
(stop when the training loss stops improving by `tol`)
It is warm-started from a stagewise fit and currently supports squared error only. Structure re-discovery makes the sweep a heuristic (no monotonicity guarantee), so the best-loss iterate is kept.
When does it actually help? Backfitting starts from the stagewise+refit fit
and keeps the best-loss iterate, so it is never worse than
StagewiseSymbolicRegressor(refit_coefficients=True) on the training data —
the only question is whether the sweeps improve on it. That hinges entirely on
whether re-discovery changes the set of selected basis functions:
Generous per-term budget (
max_complexity≥ 2–3): greedy usually already finds a sufficient basis set, so the joint least-squares refit makes the two essentially identical. Backfitting adds nothing here — prefer the stagewise regressor.Small per-term budget and collinear features (
max_complexity=1, the GAM-style single-basis regime): greedy forward selection can lock into a suboptimal basis set that a single forward pass cannot undo. Backfitting’s coordinate-descent re-discovery escapes it, changing the basis union and improving the fit — we have measured up to roughly +0.04 train / +0.06 test R² in this regime, with no downside in the cases where it does not help.
So reach for backfitting when you want small, revisable single-basis terms over correlated features (or a fixed-size GAM-style decomposition); use the stagewise regressor for larger per-term expressions. Its other forward-looking value is as the foundation for a future Bayesian backfitting variant (BART/iBART-style), which would sample a posterior over symbolic structure — genuinely beyond point-estimate SR — using the same partial-residual sweep with conjugate marginal likelihoods.
Recursive basis expansion (experimental)#
The Scope section notes that a fixed
library cannot discover compositional forms like x0*sin(x1) or exp(x0*x1).
RecursiveSymbolicRegressor is an experimental step past that ceiling. Instead
of enumerating a huge composition space up front (which explodes
combinatorially), it grows the library lazily along the residual:
Fit a sparse model over the current library; take the residual.
Compose the currently useful terms (selected terms + features) with a small operator set (unary functions, products, ratios) — one new layer.
Screen: drop non-finite candidates, deduplicate, keep the top few by correlation with the residual.
Add the survivors and refit. Repeat. The effective composition depth is the number of rounds, because a term found in one round feeds the next.
from jaxsr.additive import RecursiveSymbolicRegressor
model = RecursiveSymbolicRegressor(n_expansions=3, max_terms=6, beam_width=25)
model.fit(X, y)
print(model.expression_) # e.g. recovers "exp((x0)*(x1))" exactly
print(model.history_) # library size / n_terms / train R^2 per round
This is essentially Fast Function Extraction (FFX) / symbolic feature
construction — a deterministic, bounded cousin of genetic programming. On simple
compositional targets it substantially beats a flat library (e.g. exp(x0*x1):
R² 1.00 vs 0.70) and is competitive with a strong GP engine (matched Operon on
x0*sin(x1) in our tests).
Caveats. It re-enters search-based territory: cost grows with beam_width,
n_expansions, and feature count, and it will not match a mature GP (PySR,
Operon) on hard, high-dimensional, or deeply nested targets. The result is an
ordinary SymbolicRegressor over the grown library (so predict/expression_/
scoring and the base regressor’s non-finite-basis guard and negligible-term
pruning all apply), but the composed bases are Python closures, so the fitted
model is not serialisable via save/load, and to_sympy may not parse
deeply nested term names.