Core concepts¶
jaxfolio is small on the surface — one function shape, one result type — and deliberately uniform underneath. This page explains the architecture so the rest of the documentation reads as variations on a theme rather than a catalog of unrelated tools.
One interface¶
Every optimizer, built-in or custom, has the same signature:
returns is a wide returns panel (rows = periods, columns = assets) as a pandas
DataFrame or a 2-D array; everything after it is method-specific configuration
with sensible defaults. The return value is always a
PortfolioResult.
Because the shape is uniform, methods are interchangeable at the call site — the
backtester, compare,
the registry, and the
plots all accept any callable of this shape,
whether it ships with the library or you wrote it this morning.
flowchart LR
A[returns panel] --> B[moments<br/>μ, Σ]
B --> C{optimizer}
C --> D[PortfolioResult<br/>weights + diagnostics]
D --> E[backtest]
D --> F[plots]
D --> G[options overlay]
The moment pipeline¶
Optimizers do not consume raw returns directly; they consume moments. The
moments helper normalizes any input into a canonical
tuple:
from jaxfolio import toolkit as tk
mu, cov, names, matrix = tk.moments(returns)
# μ Σ asset T × N
# mean cov names return matrix
as_matrixaccepts aDataFrame(columns become asset names) or a raw array (names default toasset_0…), and replaces NaNs with zeros.mean_returnsandsample_covarianceare the defaults, but you can pass any covariance estimator — the library ships EWMA and Ledoit–Wolf shrinkage alongside the sample estimator.
All moment estimators return JAX arrays, so the whole chain — estimation → optimization → diagnostics — can be JIT-compiled and differentiated.
The PortfolioResult¶
Every method returns the same container:
@dataclass
class PortfolioResult:
weights: np.ndarray # sums to 1
assets: list[str] # aligned with weights
method: str # human-readable name
expected_return: float | None # annualized
volatility: float | None # annualized
sharpe: float | None # annualized
metadata: dict # free-form diagnostics
Useful helpers:
result.as_dict() # {asset: weight}, sorted by |weight| descending
result.top(5) # the 5 largest holdings by |weight|
result.weights # the raw numpy vector
The annualized expected_return, volatility, and sharpe are computed by a
single function — finalize_result — so every
method reports them on the same basis (252 periods per year by default). This is
the single source of truth for diagnostics across built-in and custom strategies
alike, which is what makes cross-method comparison fair.
The metadata dict is where method-specific detail lives: solver iteration
counts, ERC risk contributions, HRP cluster labels and linkage, Black–Litterman
posterior returns, LLM views and confidences, and so on. The plots read from it
(for example, plot_dendrogram needs the HRP
linkage).
The shared solver¶
This is the heart of the design. Every constrained classical objective — min-variance, mean-variance, max-Sharpe, max-diversification, CVaR, Black–Litterman's posterior optimization — is minimized by the same routine:
Take a gradient step on an unconstrained objective f(w), then project the
weights back onto the feasible set \(\mathcal{C}\). The default solver ("spg")
uses a spectral projected-gradient method with Barzilai–Borwein step sizes — it
tunes its own step from the local curvature and converges to the same optimum a
dedicated QP solver finds, in a few hundred iterations and with no learning-rate
to set. (An Adam solver is also available via
OptimizerConfig(solver="adam"), preferred when differentiating through the
optimizer.) The loop runs inside jax.lax.while_loop, so the entire solve
compiles to a single JIT kernel — cached per problem shape, so repeated solves
(e.g. every rebalance of a backtest) reuse it:
from jaxfolio import toolkit as tk
projection = tk.make_projection(long_only=True, weight_bounds=(0.0, 1.0))
def objective(w):
return w @ cov @ w # e.g. minimum variance
w, info = tk.solve_projected_gradient(objective, tk.equal_start(n), projection)
Only the objective changes between methods. This is why adding a new constrained optimizer is a one-line objective (see Custom strategies), and why every method inherits the same convergence behavior and constraint handling.
Closed forms where they exist
A few methods sidestep the iterative solver. Equal-weight and inverse-volatility are closed-form. Risk parity uses the provably convergent cyclical coordinate descent of Griveau-Billion, Richard & Roncalli (2013). The graph methods (HRP, HERC, MST) are combinatorial and stay in NumPy/SciPy.
Constraints as projections¶
The feasible set \(\mathcal{C}\) is chosen by
make_projection:
- Long-only, fully invested → Euclidean projection onto the probability simplex \(\{w : w \ge 0,\ \sum w = 1\}\) via the exact Duchi et al. (2008) sort algorithm.
- Bounded / shorting allowed → projection onto a box \([lo, hi]\) intersected with the budget hyperplane \(\sum w = \text{budget}\), solved by bisection on the budget multiplier.
Both projections are pure JAX and safe under jit/vmap.
OptimizerConfig¶
Constrained methods take an optional
OptimizerConfig that
controls the constraint set and solver behavior:
from jaxfolio import OptimizerConfig
cfg = OptimizerConfig(
long_only=True, # simplex projection
weight_bounds=(0.0, 0.2), # cap any single position at 20%
risk_free_rate=0.0, # per-period, for Sharpe-style objectives
l2_reg=1e-3, # optional diversification penalty
max_iter=2000,
solver="spg", # spectral projected gradient (default); or "adam"
learning_rate=None, # None = auto step from local curvature
tol=1e-7, # projected-gradient (KKT) stationarity tolerance
)
jf.maximum_sharpe(returns, config=cfg)
Differentiability, end to end¶
Because moment estimation, the solver, and the diagnostics are all JAX, the whole allocation is a differentiable function of its inputs. That unlocks three things ordinary libraries cannot do:
- Train allocation policies — differentiate through an optimizer to learn
a mapping from market state to weights (this is exactly what
deep_sharpedoes). - Exact Greeks for free — the options layer never
hand-codes a Greek; every sensitivity is
jax.gradof the same pricing function, so they can never drift out of sync. - Fast backtests — thousands of rebalances run through JIT-compiled kernels.
The registry¶
Strategies can be registered by name so they are discoverable and mixable with the built-ins:
jf.list_strategies() # every registered method, built-in + custom
jf.get_strategy("maximum_sharpe") # look up by name
jf.strategy_info("risk_parity") # name, callable, description, builtin flag
The built-in optimizers register themselves at import time, and your own
custom strategies can join them with a single
decorator or register=True.
Module map¶
| Module | Responsibility |
|---|---|
jaxfolio.optimizers |
the sixteen methods (classical, learning, graph) |
jaxfolio.moments |
mean & covariance estimators (sample, EWMA, Ledoit–Wolf) |
jaxfolio.constraints |
simplex / box-budget projections |
jaxfolio.results |
finalize_result, moments, solver-agnostic assembly |
jaxfolio.toolkit |
the public building blocks for authoring strategies |
jaxfolio.backtest |
walk-forward engine, compare, metrics |
jaxfolio.options |
pricing, autodiff Greeks, multi-leg strategies, overlays |
jaxfolio.llm |
local-model views → Black–Litterman |
jaxfolio.viz |
dark-themed matplotlib plots |
jaxfolio.data |
synthetic generator + CSV / Parquet / Yahoo loaders |
jaxfolio.registry |
name-based strategy registry |