Solvers and Utilities#
This document covers numerical solvers, uncertainty propagation, and utility functions in Difflow.
Table of Contents#
-
Equation-Oriented Solver — simultaneous solution of all unit equations
External Solvers: pounce and discopt — a flowsheet as a flat NLP or as an implicit residual block
Numerical Solvers#
External Libraries: Difflow uses optimistix for fixed-point iteration and root-finding, and diffrax for ODE integration.
All solvers in Difflow are:
Implemented using JAX primitives for automatic differentiation
Support implicit differentiation for efficient backward passes
Fully JIT-compilable for performance
Fixed-Point Iteration#
Solves equations of the form \(x^* = f(x^*, \text{args})\).
import optimistix as optx
def solve_recycle(fresh_feed):
"""Fixed-point iteration for recycle streams."""
def flowsheet_iteration(recycle, args):
"""Update function: recycle_out = f(recycle_in)"""
# Mix fresh feed with recycle, run process, return new recycle
...
return new_recycle
# Create solver with tolerances
solver = optx.FixedPointIteration(rtol=1e-8, atol=1e-8)
# Solve for steady state
solution = optx.fixed_point(
fn=flowsheet_iteration,
solver=solver,
y0=initial_guess,
args=fresh_feed,
max_steps=100,
throw=False # Return solution even if not fully converged
)
return solution.value
Algorithm:
Where \(\alpha\) is the damping factor.
Convergence Criterion:
Example: Solving recycle loop
import jax.numpy as jnp
import optimistix as optx
fresh_feed = 1.0 # mol/s of A
reactor_params = {'conversion': 0.6, 'recycle_fraction': 0.9}
def recycle_update(recycle_flow, args):
feed, p = args
# Mix fresh feed with recycle
mixed = feed + recycle_flow
# Reactor: unconverted A leaves, a fraction of it is recycled
unconverted = mixed * (1.0 - p['conversion'])
# Return new recycle flow (converges when input = output)
return p['recycle_fraction'] * unconverted
# Solve for steady-state recycle flow
solver = optx.FixedPointIteration(rtol=1e-6, atol=1e-6)
solution = optx.fixed_point(
fn=recycle_update,
solver=solver,
y0=jnp.asarray(0.1), # Initial guess
args=(fresh_feed, reactor_params),
max_steps=100
)
recycle_ss = solution.value
Newton-Raphson Solver#
Solves equations of the form \(g(x^*, \text{args}) = 0\).
import optimistix as optx
def solve_equation(args):
"""Newton-Raphson root finding."""
def residual_function(x, args):
"""Function whose root we seek: g(x, args) = 0"""
# Return residual that should equal zero
return ...
# Create Newton solver
solver = optx.Newton(rtol=1e-10, atol=1e-10)
# Find root
solution = optx.root_find(
fn=residual_function,
solver=solver,
y0=initial_guess,
args=args,
max_steps=50
)
return solution.value
Algorithm:
Where \(J = \frac{\partial g}{\partial x}\) is the Jacobian, computed automatically using JAX.
Features:
Automatic Jacobian computation via
jax.jacobianCustom VJP for efficient backward differentiation
Line search for improved robustness (optional)
Example: Bubble point temperature
import jax.numpy as jnp
import optimistix as optx
# Mole fractions, pressure and a simple K-value model (illustrative constants)
liquid_composition = jnp.array([0.4, 0.6])
pressure = 101325.0
A, B = jnp.array([10.8, 11.5]), jnp.array([2800.0, 3700.0])
def K_values_func(T):
return jnp.exp(A - B / T) * 1e5 / pressure
def bubble_residual(T, args):
x, P, K_func = args
K = K_func(T)
# Bubble point: sum(x_i * K_i) = 1
return jnp.sum(x * K) - 1.0
solver = optx.Newton(rtol=1e-8, atol=1e-8)
solution = optx.root_find(
fn=bubble_residual,
solver=solver,
y0=jnp.asarray(350.0),
args=(liquid_composition, pressure, K_values_func),
max_steps=50
)
T_bubble = solution.value
Rachford-Rice Solver#
Specialized solver for flash calculations. Uses optimistix’s bisection method with automatic bounds.
import optimistix as optx
def rachford_rice(psi, z, K):
"""Rachford-Rice equation: should equal zero at solution."""
return jnp.sum(z * (K - 1) / (1 + psi * (K - 1)))
def solve_flash(z, K):
"""Solve for vapor fraction using bisection."""
# Bisection bounds for vapor fraction
psi_min = 1 / (1 - jnp.max(K)) + 1e-6
psi_max = 1 / (1 - jnp.min(K)) - 1e-6
psi_min = jnp.clip(psi_min, 0.0, None)
psi_max = jnp.clip(psi_max, None, 1.0)
solver = optx.Bisection(rtol=1e-10, atol=1e-10)
solution = optx.root_find(
fn=rachford_rice,
solver=solver,
y0=0.5,
args=(z, K),
options={"lower": psi_min, "upper": psi_max},
max_steps=50
)
return solution.value
# Get phase compositions from vapor fraction
def phase_compositions(z, K, psi):
x = z / (1 + psi * (K - 1)) # Liquid
y = K * x # Vapor
return x, y
Rachford-Rice Equation:
Algorithm:
Bound vapor fraction: \(V \in [V_{min}, V_{max}]\)
\(V_{min} = \max_i \frac{K_i z_i - 1}{K_i - 1}\) for \(K_i > 1\)
\(V_{max} = \min_i \frac{1 - z_i}{1 - K_i}\) for \(K_i < 1\)
Newton iteration with bounds enforcement
Damping for stability near boundaries
Phase Compositions:
ODE Integration#
Used internally by PFR, fed-batch reactor, and other dynamic models. Difflow uses diffrax for ODE integration.
import diffrax
def integrate_ode(y0, t_span, derivative_fn, args):
"""Integrate ODE using diffrax."""
def vector_field(t, y, args):
"""dy/dt = f(t, y, args)"""
return derivative_fn(y, args)
# Define the ODE term
term = diffrax.ODETerm(vector_field)
# Choose a solver (adaptive)
solver = diffrax.Tsit5() # 5th order adaptive method
# Configure step size controller
stepsize_controller = diffrax.PIDController(rtol=1e-6, atol=1e-8)
# Optionally save at specific points
saveat = diffrax.SaveAt(ts=jnp.linspace(t_span[0], t_span[1], 101))
# Solve
solution = diffrax.diffeqsolve(
term,
solver,
t0=t_span[0],
t1=t_span[1],
dt0=0.01, # Initial step size
y0=y0,
args=args,
saveat=saveat,
stepsize_controller=stepsize_controller
)
return solution.ys # Trajectory at saved points
Why diffrax?
Adaptive step size control for efficiency and accuracy
Fully differentiable through the integration (supports adjoint methods)
JIT-compilable
Multiple solvers available (Tsit5, Dopri5, Kvaerno5 for stiff problems)
Uncertainty Propagation#
Location: difflow/uncertainty.py
Linear Propagation#
First-order Taylor expansion for uncertainty propagation.
import jax
import jax.numpy as jnp
from difflow.uncertainty import linear_propagation
# Define model and uncertainties: models take a dict of parameters
def model(params):
# First-order reaction in a CSTR: conversion = k*tau / (1 + k*tau)
k = 1e5 * jnp.exp(-6000.0 / params['T'])
tau = 10.0 / params['F']
return k * tau / (1.0 + k * tau)
nominal = {'T': 350.0, 'F': 1.0}
uncertainties = {'T': 5.0, 'F': 0.05} # 1-sigma
# Propagate uncertainty (returns output, std and a sensitivity dict)
mean, std, info = linear_propagation(model, nominal, uncertainties)
print(f"Conversion: {mean:.4f} +/- {std:.4f}")
Theory:
For \(y = f(\mathbf{x})\) with \(\mathbf{x} \sim N(\boldsymbol{\mu}, \boldsymbol{\Sigma})\):
Where \(\mathbf{J} = \nabla f|_{\boldsymbol{\mu}}\) is the Jacobian.
For uncorrelated inputs (\(\boldsymbol{\Sigma}\) is diagonal):
Advantages:
Fast (single gradient evaluation)
Accurate for small uncertainties and linear systems
Limitations:
First-order approximation
May underestimate uncertainty for nonlinear systems
Monte Carlo Propagation#
Sampling-based uncertainty propagation.
from difflow.uncertainty import monte_carlo_propagation
# Monte Carlo analysis
mean, std, info = monte_carlo_propagation(
model=model,
nominal_params=nominal,
uncertainties=uncertainties,
n_samples=10000,
distribution='normal' # or 'uniform'
)
print(f"Mean: {float(mean):.4f}")
print(f"Std: {float(std):.4f}")
print(f"2.5th percentile: {info['p2.5']:.4f}")
print(f"97.5th percentile: {info['p97.5']:.4f}")
Algorithm:
Generate \(N\) samples from input distribution
Evaluate model for each sample (vectorized with
vmap)Compute output statistics
Implementation (vectorized for efficiency):
import jax.numpy as jnp
from jax import vmap, random
def monte_carlo_sketch(model, nominal, uncertainties, n_samples, key=None):
if key is None:
key = random.PRNGKey(0)
# Generate samples
samples = nominal + uncertainties * random.normal(key, shape=(n_samples, len(nominal)))
# Vectorized model evaluation (model takes an array here)
outputs = vmap(model)(samples)
return {
'mean': jnp.mean(outputs),
'std': jnp.std(outputs),
'p5': jnp.percentile(outputs, 5),
'p95': jnp.percentile(outputs, 95),
'samples': outputs
}
Advantages:
Accurate for nonlinear systems
Provides full distribution, not just mean/variance
Handles non-Gaussian distributions
Sensitivity Analysis#
Gradient-based sensitivity analysis.
from difflow.uncertainty import sensitivity_analysis
# Compute sensitivities
sensitivities = sensitivity_analysis(
model=model,
nominal_params=nominal,
param_ranges={'T': (340.0, 360.0), 'F': (0.8, 1.2)}, # optional; default +/-20%
)
for name, sens in sensitivities.items():
print(f"{name}: gradient = {float(sens['gradient']):.4f}, elasticity = {float(sens['elasticity']):.4f}")
Sensitivity Metrics:
Gradient (absolute sensitivity): $\(S_i = \frac{\partial y}{\partial x_i}\)$
Normalized sensitivity (dimensionless): $\(S_i^* = \frac{\partial y}{\partial x_i} \cdot \frac{x_i}{y}\)$
Variance contribution: $\(C_i = \frac{S_i^2 \sigma_{x_i}^2}{\sum_j S_j^2 \sigma_{x_j}^2}\)$
Sobol Indices#
Global sensitivity analysis using Sobol variance decomposition.
from difflow.uncertainty import sobol_indices
# Compute Sobol indices
indices = sobol_indices(
model=model,
param_bounds={'T': (340.0, 360.0), 'F': (0.8, 1.2)},
n_samples=1024,
)
for name, idx in indices.items():
print(f"{name}: S1 = {idx['S1']:.3f} (normalized {idx['S1_normalized']:.3f})")
sobol_indices returns first-order indices only (S1, S1_normalized,
bounds); the total-order index below is the definition, not something the
function computes. For production-grade total-order indices use a dedicated
package such as SALib.
Theory:
Total variance decomposition:
First-order Sobol index (main effect):
Total-order Sobol index (main + interactions):
Interpretation:
\(S_i \approx S_{Ti}\): Parameter has mostly main effects
\(S_{Ti} \gg S_i\): Parameter has significant interactions
\(\sum_i S_{Ti} \approx 1\): Weak interactions
\(\sum_i S_{Ti} \gg 1\): Strong interactions
Covariance Propagation#
General covariance matrix propagation.
from difflow.uncertainty import propagate_covariance
# Full covariance matrix (correlated inputs)
cov_input = jnp.array([
[25.0, 2.0], # Var(T) = 25, Cov(T,F) = 2
[2.0, 0.01] # Cov(F,T) = 2, Var(F) = 0.01
])
# Propagate covariance (the Jacobian is computed internally)
output, cov_output, jacobian = propagate_covariance(
model, nominal, cov_input, param_order=['T', 'F'])
Equation:
Implicit Differentiation#
Difflow uses implicit differentiation to compute gradients through iterative solvers.
Theory#
For a solution \(x^* = f(x^*, \theta)\) (fixed-point) or \(g(x^*, \theta) = 0\) (root-finding), the gradient w.r.t. parameters \(\theta\) is:
Fixed-point: $\(\frac{dx^*}{d\theta} = \left(I - \frac{\partial f}{\partial x}\bigg|_{x^*}\right)^{-1} \frac{\partial f}{\partial \theta}\bigg|_{x^*}\)$
Root-finding: $\(\frac{dx^*}{d\theta} = -\left(\frac{\partial g}{\partial x}\bigg|_{x^*}\right)^{-1} \frac{\partial g}{\partial \theta}\bigg|_{x^*}\)$
Implementation with Optimistix#
Optimistix handles implicit differentiation automatically. When you use optx.fixed_point() or optx.root_find(), gradients are computed using the implicit function theorem rather than by backpropagating through the solver iterations.
import optimistix as optx
import jax
import jax.numpy as jnp
def optimize_with_gradients(params):
"""Example showing automatic gradient computation through solver."""
def my_fixed_point(x, args):
# Fixed-point function that depends on params
return params['a'] * jnp.cos(x) + args
solver = optx.FixedPointIteration(rtol=1e-8, atol=1e-8)
solution = optx.fixed_point(
fn=my_fixed_point,
solver=solver,
y0=jnp.asarray(0.5),
args=0.1,
max_steps=100
)
# The solution is differentiable w.r.t. params
return solution.value ** 2
# Gradients computed via implicit differentiation
grad_fn = jax.grad(optimize_with_gradients)
params = {'a': jnp.asarray(0.4)}
gradients = grad_fn(params)
Note on Numerical Challenges: Implicit differentiation requires inverting a matrix \((I - \partial f/\partial x)\) at the solution. This can become singular or ill-conditioned when:
The system is near a bifurcation point
Reactions go to very high conversion (near 100%)
The system is at a phase boundary
Recycle ratios are extreme
If gradient computation fails while the forward solve succeeds, consider using finite differences for gradients or reformulating the problem.
Benefits#
Memory efficient: Only stores solution, not iteration history
Accurate gradients: Uses implicit function theorem, not unrolling
Fast backward pass: Single linear solve instead of backprop through iterations
Utility Functions#
Numerical Helpers#
import jax.numpy as jnp
from difflow.numerics import (
safe_divide,
safe_log,
safe_sqrt,
safe_exp,
smooth_max,
smooth_min,
smooth_clamp,
)
a, b = jnp.array(2.0), jnp.array(0.0)
# Safe operations (avoid NaN/Inf)
x = safe_divide(a, b) # Uses eps (1e-10) in place of a zero denominator
y = safe_log(jnp.array(0.0)) # Clips x to avoid log(0)
z = safe_sqrt(jnp.array(-1.0)) # Clips x to avoid sqrt(negative)
# Smooth approximations (differentiable)
max_val = smooth_max(a, b, alpha=10.0) # Log-sum-exp approximation
min_val = smooth_min(a, b, alpha=10.0) # Log-sum-exp approximation
Smooth Approximations#
For optimization, smooth approximations of non-differentiable functions:
Smooth maximum: $\(\text{softmax}(a, b) = \frac{a e^{\alpha a} + b e^{\alpha b}}{e^{\alpha a} + e^{\alpha b}}\)$
As \(\alpha \to \infty\), approaches \(\max(a, b)\).
Smooth absolute value: $\(|x|_\epsilon \approx \sqrt{x^2 + \epsilon^2}\)$
Smooth ReLU: $\(\text{softplus}(x) = \frac{1}{\beta} \log(1 + e^{\beta x})\)$
Unit conversions are deliberately not wrapped in helper functions: difflow
works in SI throughout (K, Pa, mol, kg, J), and thermodynamic properties come
from difflow.thermo and difflow.eos.
Best Practices#
Solver Selection#
Problem Type |
Recommended Solver |
|---|---|
Fixed-point (well-behaved) |
|
Fixed-point (difficult) |
Increase max_steps, adjust tolerances |
Root-finding |
|
Flash calculation |
|
ODE integration |
|
Convergence Tips#
Good initial guess: Use physical intuition or simpler model
Appropriate tolerance: 1e-6 to 1e-10 depending on application
Damping: Start with 0.3-0.5 for difficult problems
Bounds: Enforce physical constraints (positive concentrations, etc.)
Uncertainty Analysis Workflow#
import jax.numpy as jnp
from difflow.uncertainty import (
linear_propagation, sensitivity_analysis, monte_carlo_propagation, sobol_indices,
)
# 1. Define model
def process_model(params):
# ... process simulation ...
return params['F'] * jnp.exp(-1000.0 / params['T']) * (params['P'] / 101325.0)
# 2. Identify uncertain parameters
nominal = {'T': 350.0, 'P': 101325.0, 'F': 1.0}
uncertainties = {'T': 10.0, 'P': 5000.0, 'F': 0.1}
# 3. Quick screening with linear propagation
mean, std, info = linear_propagation(process_model, nominal, uncertainties)
# 4. Identify important parameters with sensitivity analysis
sens = sensitivity_analysis(process_model, nominal)
# 5. Detailed analysis on key parameters with Monte Carlo
mc_mean, mc_std, mc_info = monte_carlo_propagation(
process_model, nominal, uncertainties, n_samples=10000)
# 6. Global sensitivity with Sobol indices (if needed)
bounds = {k: (v - 2 * uncertainties[k], v + 2 * uncertainties[k]) for k, v in nominal.items()}
sobol = sobol_indices(process_model, bounds, n_samples=1024)
Debugging Numerical Issues#
# Check for NaN/Inf
import jax.numpy as jnp
def check_numerics(x, name="value"):
if jnp.any(jnp.isnan(x)):
print(f"NaN detected in {name}")
if jnp.any(jnp.isinf(x)):
print(f"Inf detected in {name}")
return x
# Monitor convergence
def solve_with_monitoring(f, x0, args, tol, max_iter):
x = x0
for i in range(max_iter):
x_new = f(x, args)
error = jnp.max(jnp.abs(x_new - x))
print(f"Iter {i}: error = {error:.2e}")
if error < tol:
print(f"Converged in {i+1} iterations")
return x_new
x = x_new
print("Warning: Did not converge")
return x