Parameter Estimation#
The difflow.estimation module provides a JAX-powered parameter estimation
framework inspired by pyomo.parmest. It combines automatic differentiation with
structured APIs for model fitting, uncertainty quantification, and diagnostics, so
you can calibrate models to experimental data and propagate the resulting
uncertainty through differentiable flowsheets.
Table of Contents#
Overview#
The module fits the parameters of a user-defined model to one or more
Experiment observations, using scipy.optimize driven by exact JAX
gradients (no finite differences). Because every operation is differentiable
end-to-end, the same machinery works whether you are fitting a simple algebraic
model or calibrating parameters inside a complex flowsheet with recycles.
Key capabilities:
Parameter estimation with box bounds and multiple objective functions
Fisher-information confidence intervals from the exact Hessian
Parametric and nonparametric bootstrap uncertainty
Cross-validation for predictive performance
Regression diagnostics (R², RMSE, AIC, BIC)
Multi-output and nonlinear model support
Architecture#
Core Components#
Experiment(experiment.py)Encapsulates a single experimental observation
Stores inputs, observed outputs, uncertainties, and metadata
Inherits from
ParamsMixinfor a dict-like interfaceConvenience properties:
observed_array,weights,output_names
Estimator(estimator.py)Main orchestrator class
Wraps fitting, confidence intervals, bootstrap, cross-validation, and diagnostics
Uses
scipy.optimizewith exact JAX gradientsSupports parameter bounds and multiple objective functions
EstimationResult(estimator.py)Returns fitted parameters (dict and array forms)
Reports convergence status, iterations, and objective value
Inherits from
ParamsMixin
Supporting Modules#
Objectives (
objectives.py) — all fully differentiablesum_squared_errors(SSE)weighted_sum_squared_errors(WSSE)negative_log_likelihood(NLL)
Confidence Intervals (
confidence.py)Fisher information matrix approach
Computes covariance, standard errors, and confidence intervals
Uses the JAX Hessian for exact second derivatives
Diagnostics (
diagnostics.py)R², adjusted R², RMSE
AIC and BIC information criteria
Residuals analysis
Bootstrap (
bootstrap.py)Nonparametric (resample experiments)
Parametric (resample residuals)
Percentile confidence intervals
Takes an
objective, and it must be the onefitminimized (see Matching the objective)
Cross-Validation (
cross_validation.py)Leave-N-out cross-validation
Predictive performance metrics
Optional
max_foldslimit for large datasets
Identifiability (
identifiability.py)check_identifiability— rank test on the sensitivity matrixAnswers whether the parameters are separately estimable at all
Reuses the SVD rank machinery of
difflow.reconciliation.structure
Experiment Design (
design.py)design_experiments— which runs to do next (D/A/E/modified-E optimal)predicted_covariance— the confidence intervals a campaign would buySelects from a candidate list, for any model JAX can differentiate. Continuous design optimization, profile likelihood, model discrimination and classical designs are in the separate
discopt-doeplugin (discopt.doe), which needs adiscopt.modelingmodel rather than a JAX one.
Basic Workflow#
Before fitting anything, ask whether the parameters can be told apart at all:
check_identifiability runs a rank test on the sensitivity matrix, and when it
fails no estimator on this page can help — see
Experiment Design and Identifiability, which also covers
choosing the next experiment.
The fitting API then follows a simple fit → quantify → diagnose pattern:
import jax.numpy as jnp
from difflow.estimation import Estimator, Experiment
# model_fn(theta, experiment) -> predicted outputs (must be JAX-differentiable)
def model_fn(theta, exp):
return {'y': theta['k'] * jnp.exp(-theta['b'] * exp.inputs['t'])}
param_names = ['k', 'b']
param_bounds = {'k': (0.0, 10.0), 'b': (0.0, 5.0)}
theta_init = {'k': 1.0, 'b': 0.5}
experiments = [
Experiment(inputs={'t': t}, observed={'y': 2.0 * float(jnp.exp(-0.7 * t)) + 0.01 * (-1) ** i})
for i, t in enumerate([0.0, 0.5, 1.0, 1.5, 2.0, 3.0])
]
est = Estimator(model_fn, param_names, param_bounds)
result = est.fit(experiments, theta_init) # fit parameters
ci = est.confidence_intervals(result, experiments) # Fisher-information CIs
diag = est.diagnostics(result, experiments) # R², RMSE, AIC, BIC
bs = est.bootstrap(result, experiments, n_bootstrap=20) # bootstrap uncertainty
print(est.summary(result, experiments))
result exposes the fitted parameters as both a dict and an array, along with the
convergence status and final objective value.
Matching the objective#
fit, confidence_intervals, bootstrap and summary each take an objective,
and all four default to 'sse'. A weighted fit must carry objective='wsse'
through every one of them. The default is kept for backwards compatibility, and
it is silently wrong in exactly the case where weighting mattered:
result = est.fit(experiments, theta_init, objective='wsse')
ci = est.confidence_intervals(result, experiments, objective='wsse')
bs = est.bootstrap(result, experiments, n_bootstrap=20, objective='wsse')
print(est.summary(result, experiments, objective='wsse'))
The intervals come from the Hessian of the objective at the optimum, so an
unweighted Hessian over a weighted fit reports the standard errors of numbers
nobody computed. The bootstrap is worse than that: it refits, so on the default
it returns the sampling distribution of a different estimator. In
examples/22_ree_parameter_estimation.ipynb, where the measured concentrations
span four decades, the mismatch inflates the bootstrap standard errors five-fold
and reads as a failed linearization.
Uncertainty Quantification#
The module offers three complementary approaches, each with different assumptions and cost:
Fisher information — asymptotic and fast; derived from the exact Hessian of the objective at the optimum. Best when the model is approximately linear near the solution.
Bootstrap — distribution-free and robust; resamples experiments (nonparametric) or residuals (parametric) and refits. More expensive but makes fewer assumptions.
Cross-validation — assesses predictive performance rather than parameter precision, and helps detect over-fitting.
Using more than one gives complementary insight into how well the parameters are determined.
Use in Chemical Engineering#
The module is well suited to calibrating process models against measured data:
Kinetics — fit Arrhenius parameters (A, Ea) from batch-reactor concentration profiles; estimate catalyst deactivation kinetics.
Thermodynamics — calibrate distribution coefficients, equilibrium constants, and activity-coefficient model parameters.
Transport — estimate mass- and heat-transfer coefficients from measured profiles.
Model calibration — fit flowsheet models and unit-operation efficiency factors to plant data.
Because difflow units are already differentiable, estimated parameters can be
fit inside complete flowsheets — including those with recycles — without any extra
machinery.
Design Principle: Fit to Measured Quantities#
Fit to the quantities you actually measure, not to derived values.
✅ Good: fit to concentrations, temperatures, pressures.
❌ Avoid: fitting to derived quantities such as distribution coefficients
D = C_org / C_aq, reaction rates, or heat-transfer rates.
Fitting raw measurements:
handles error propagation correctly,
provides more information (e.g. two measured concentrations rather than one ratio),
matches the actual experimental workflow, and
enables mass-balance and other physical constraints.
The example notebook demonstrates this by fitting equilibrium concentrations
(C_aq, C_org) directly rather than the distribution coefficients derived from
them.
Current Limitations#
The current implementation prioritizes clarity and correctness. Known limitations to be aware of:
Optimization uses
scipy.optimize.minimize; JAX-native optimizers (optax,jaxopt) are not yet wired in.Constraints are limited to parameter box bounds (no linear/nonlinear constraints).
Weighting supports inverse-variance only (no robust M-estimators or correlated errors).
Bootstrap runs sequentially, so large resample counts can be slow.
Identifiability diagnostics are local and rank-based (
check_identifiability, see Experiment Design and Identifiability); there is no profile likelihood and no global symbolic identifiability analysis.There is no model discrimination, no estimability ranking, no classical or screening design, and no continuous optimization of experimental conditions. Those live in the
discopt-doeplugin; see what is here and what is indiscopt-doefor the division of labour and when to reach for which.
Example#
A complete worked example is available in the Examples section: REE parameter estimation, which recovers the shipped PC88A coefficients for La, Nd and Dy from measured liquid–liquid extraction concentrations. It runs the identifiability check first, then fits a ladder of four models — free quadratic, linear, shared slope, and the one-parameter mass-action form the database uses — and lets AIC and BIC choose, with Fisher and bootstrap uncertainty and censoring of below-LOQ measurements.