Why Differentiable Flowsheets?#
Prerequisites: All previous tutorials (00a-00j)
Learning Objectives:
Understand why gradients through flowsheets are powerful
See how difflow differentiates through recycle iterations
Explore applications: optimization, sensitivity analysis, uncertainty propagation
Connect to the advanced examples in the repository
The Big Picture#
You’ve learned how to:
Represent flowsheets as systems of equations
Solve single units (CSTR, PFR, flash)
Connect units in series and parallel
Handle recycles with fixed-point iteration
Now comes the payoff: Because everything in difflow is built with JAX, we can compute exact gradients through the entire flowsheet!
Why Gradients Matter#
Traditional Approach: Finite Differences#
To compute sensitivity of output \(y\) to parameter \(\theta\):
Problems:
Need to re-solve flowsheet for each parameter → slow
Numerical errors from finite \(\epsilon\)
Scales poorly with many parameters
Difflow Approach: Automatic Differentiation#
JAX computes exact gradients through:
All unit operations
All thermodynamic calculations
Even through recycle iterations!
Benefits:
One backward pass gives ALL gradients
Machine precision (no numerical errors)
Enables gradient-based optimization
# Setup
import jax.numpy as jnp
import jax
jax.config.update("jax_enable_x64", True)
import matplotlib.pyplot as plt
import numpy as np
import optimistix as optx
from difflow import (CSTR, CSTRParams, Flash, FlashParams,
make_stream, get_flows, IdealThermo, SpeciesData)
from difflow.units.flash import Mixer
WARNING:2026-01-10 21:13:08,243:jax._src.xla_bridge:852: An NVIDIA GPU may be present on this machine, but a CUDA-enabled jaxlib is not installed. Falling back to cpu.
# Set up the reactor-separator-recycle flowsheet from tutorial 00j
species_data = {
'A': SpeciesData(name='A', MW=100.0, Cp_coeffs=(100.0, 0, 0, 0),
Hvap_coeffs=(40000.0, 0.38, 500.0),
antoine_coeffs=(10.0, 2200.0, -40.0), Hf=0.0),
'B': SpeciesData(name='B', MW=80.0, Cp_coeffs=(80.0, 0, 0, 0),
Hvap_coeffs=(32000.0, 0.38, 450.0),
antoine_coeffs=(10.0, 1600.0, -40.0), Hf=-50000.0),
}
thermo = IdealThermo(species_data)
species_order = ['A', 'B']
def rate_fn(C, T, params):
return jnp.array([params['k'] * C['A']])
stoich = jnp.array([[-1.0], [1.0]])
def solve_flowsheet(params, fresh_feed):
"""
Solve reactor-separator-recycle flowsheet with purge.
Returns product B flow.
The purge stream removes some recycle material to prevent accumulation
and makes the system more sensitive to operating parameters.
"""
V_reactor = params['V_reactor']
k = params['k']
T_reactor = params['T_reactor']
T_flash = params['T_flash']
P_flash = params['P_flash']
Q_vol = params['Q_vol']
purge_frac = params.get('purge_frac', jnp.array(0.1)) # 10% purge
# Create units
cstr_params = CSTRParams(
V=V_reactor,
rate_fn=rate_fn,
stoich=stoich,
rate_params={'k': k},
species_order=species_order,
)
reactor = CSTR(cstr_params, thermo=thermo, mode='isothermal')
flash_params = FlashParams(species_order=species_order)
flash = Flash(flash_params, thermo=thermo)
mixer = Mixer(species_order, thermo=thermo)
def flowsheet_step(recycle_arr, args):
fresh = args
recycle = make_stream({'A': recycle_arr[0], 'B': recycle_arr[1]}, T=T_flash, P=P_flash)
inlet, _ = mixer(fresh, recycle)
out, _ = reactor(inlet, T_spec=T_reactor, volumetric_flow=Q_vol)
liquid, vapor, _ = flash(out, T=T_flash, P=P_flash)
# Apply purge: only (1-purge_frac) of liquid recycles
recycle_A = liquid['F_A'] * (1 - purge_frac)
recycle_B = liquid['F_B'] * (1 - purge_frac)
return jnp.array([recycle_A, recycle_B])
# Solve recycle using optimistix with more iterations
solver = optx.FixedPointIteration(rtol=1e-6, atol=1e-6)
solution = optx.fixed_point(flowsheet_step, solver, jnp.array([1.0, 0.1]),
args=fresh_feed, max_steps=500, throw=False)
converged = solution.value
# Final evaluation
recycle = make_stream({'A': converged[0], 'B': converged[1]}, T=T_flash, P=P_flash)
inlet, _ = mixer(fresh_feed, recycle)
out, _ = reactor(inlet, T_spec=T_reactor, volumetric_flow=Q_vol)
liquid, vapor, _ = flash(out, T=T_flash, P=P_flash)
return vapor['F_B']
# Test
params = {
'V_reactor': jnp.array(2.0),
'k': jnp.array(0.3),
'T_reactor': jnp.array(350.0),
'T_flash': jnp.array(350.0),
'P_flash': jnp.array(50000.0),
'Q_vol': jnp.array(0.1),
'purge_frac': jnp.array(0.1), # 10% purge - creates material loss
}
fresh_feed = make_stream({'A': 10.0, 'B': 0.0}, T=300.0, P=101325.0)
F_B_product = solve_flowsheet(params, fresh_feed)
print(f"Product B: {float(F_B_product):.4f} mol/s")
print(f"Overall yield: {float(F_B_product)/10*100:.1f}%")
print(f"\nNote: 10% purge reduces yield but makes sensitivities visible")
Product B: 9.3864 mol/s
Overall yield: 93.9%
Note: 10% purge reduces yield but makes sensitivities visible
Application 1: Sensitivity Analysis#
Question: How does product flow change with each parameter?
With JAX, we get ALL sensitivities in ONE call!
# Compute gradients of product B w.r.t. all parameters
grad_fn = jax.grad(lambda p: solve_flowsheet(p, fresh_feed))
grads = grad_fn(params)
print("Sensitivity Analysis: ∂(Product B)/∂(Parameter)")
print("=" * 60)
print(f"")
print(f"{'Parameter':<15} {'Value':<15} {'Sensitivity':<20} {'Interpretation'}")
print("-" * 80)
interpretations = {
'V_reactor': 'mol/s per m³',
'k': 'mol/s per (1/s)',
'T_reactor': 'mol/s per K',
'T_flash': 'mol/s per K',
'P_flash': 'mol/s per Pa',
'Q_vol': 'mol/s per (m³/s)',
'purge_frac': 'mol/s per fraction',
}
for key in params:
val = float(params[key])
grad = float(grads[key])
print(f"{key:<15} {val:<15.2f} {grad:<20.6f} {interpretations[key]}")
Sensitivity Analysis: ∂(Product B)/∂(Parameter)
============================================================
Parameter Value Sensitivity Interpretation
--------------------------------------------------------------------------------
V_reactor 2.00 0.287805 mol/s per m³
k 0.30 1.918701 mol/s per (1/s)
T_reactor 350.00 0.000000 mol/s per K
T_flash 350.00 0.053890 mol/s per K
P_flash 50000.00 -0.000028 mol/s per Pa
Q_vol 0.10 -5.756104 mol/s per (m³/s)
purge_frac 0.10 -5.607789 mol/s per fraction
# Visualize sensitivities
fig, ax = plt.subplots(figsize=(10, 5))
# Normalize sensitivities for comparison (% change in output per % change in input)
normalized_sens = {}
base_output = float(F_B_product)
# Avoid division by zero
if abs(base_output) < 1e-10:
print("Warning: Product flow is near zero. Using absolute sensitivities.")
for key in params:
normalized_sens[key] = float(grads[key])
else:
for key in params:
# Elasticity: (∂y/y) / (∂x/x) = (∂y/∂x) * (x/y)
normalized_sens[key] = float(grads[key]) * float(params[key]) / base_output
keys = list(normalized_sens.keys())
values = [normalized_sens[k] for k in keys]
colors = ['green' if v > 0 else 'red' for v in values]
ax.barh(keys, values, color=colors, alpha=0.7)
ax.axvline(x=0, color='k', linewidth=0.5)
ax.set_xlabel('Normalized Sensitivity (% change in product per % change in parameter)', fontsize=10)
ax.set_title('Which Parameters Matter Most?', fontsize=12)
ax.grid(True, alpha=0.3, axis='x')
plt.tight_layout()
print("\nInterpretation:")
print(" Green bars: Increasing parameter increases product")
print(" Red bars: Increasing parameter decreases product")
print(" Longer bars: Parameter has more influence")
Interpretation:
Green bars: Increasing parameter increases product
Red bars: Increasing parameter decreases product
Longer bars: Parameter has more influence
Using Gradients to Guide Improvement#
The sensitivity analysis told us that ∂F_B/∂T_flash = 0.054 mol/s per K — a positive value meaning higher flash temperature gives more product. Let’s verify this and see how much we can gain.
# Demonstrate the effect of increasing T_flash
# The gradient tells us the direction of improvement (at the current operating point)
T_flash_values = [350.0, 355.0, 360.0, 365.0, 370.0]
results = []
print("Effect of Flash Temperature on Product Yield")
print("=" * 60)
print(f"{'T_flash (K)':<15} {'Product B (mol/s)':<20} {'Yield (%)':<15}")
print("-" * 60)
for T in T_flash_values:
params_test = params.copy()
params_test['T_flash'] = jnp.array(T)
F_B = solve_flowsheet(params_test, fresh_feed)
yield_pct = float(F_B) / 10.0 * 100
results.append((T, float(F_B), yield_pct))
print(f"{T:<15.0f} {float(F_B):<20.4f} {yield_pct:<15.1f}")
# Find best result
best_idx = max(range(len(results)), key=lambda i: results[i][2])
best_T, best_FB, best_yield = results[best_idx]
baseline_yield = results[0][2]
# Plot the trend
fig, ax = plt.subplots(figsize=(8, 5))
temps = [r[0] for r in results]
yields = [r[2] for r in results]
ax.plot(temps, yields, 'bo-', linewidth=2, markersize=8)
ax.axhline(y=baseline_yield, color='r', linestyle='--', alpha=0.5,
label=f'Baseline (350 K): {baseline_yield:.1f}%')
ax.plot(best_T, best_yield, 'g*', markersize=15, label=f'Best ({best_T:.0f} K): {best_yield:.1f}%')
ax.set_xlabel('Flash Temperature (K)', fontsize=12)
ax.set_ylabel('Product Yield (%)', fontsize=12)
ax.set_title('Gradient Predicts Direction of Improvement', fontsize=12)
ax.legend()
ax.grid(True, alpha=0.3)
plt.tight_layout()
improvement = best_yield - baseline_yield
print(f"\nThe gradient ∂F_B/∂T_flash > 0 correctly predicts that")
print(f"increasing T_flash from 350 K improves yield.")
print(f"")
print(f"Best yield at T_flash = {best_T:.0f} K: {best_yield:.1f}% (+{improvement:.1f}% vs baseline)")
print(f"")
print(f"Note: The relationship is nonlinear - yield peaks then decreases.")
print(f"This is why optimization tools are valuable for finding the true optimum!")
Effect of Flash Temperature on Product Yield
============================================================
T_flash (K) Product B (mol/s) Yield (%)
------------------------------------------------------------
350 9.3864 93.9
355 9.5371 95.4
360 9.5750 95.8
365 9.5643 95.6
370 9.5215 95.2
The gradient ∂F_B/∂T_flash > 0 correctly predicts that
increasing T_flash from 350 K improves yield.
Best yield at T_flash = 360 K: 95.8% (+1.9% vs baseline)
Note: The relationship is nonlinear - yield peaks then decreases.
This is why optimization tools are valuable for finding the true optimum!
Application 2: Gradient-Based Optimization#
Goal: Find the reactor volume that maximizes yield while keeping cost reasonable.
With gradients, we can use efficient optimizers like gradient descent!
def objective(V, args):
"""
Objective: Maximize product - cost.
Product value: $100/mol
Reactor cost: $50/m³ (annualized)
Returns negative profit (we minimize).
"""
params_opt, fresh_feed = args
params_opt = params_opt.copy()
params_opt['V_reactor'] = V[0] # V is 1D array
F_B = solve_flowsheet(params_opt, fresh_feed)
revenue = 100.0 * F_B # $/s
cost = 50.0 * V[0] # $/s (simplified)
return -(revenue - cost) # Negative because we minimize
# Use optimistix BFGS optimizer
solver = optx.BFGS(rtol=1e-5, atol=1e-5)
# Initial guess
V_init = jnp.array([1.0])
# Optimize
solution = optx.minimise(
objective,
solver,
V_init,
args=(params, fresh_feed),
max_steps=100,
throw=False
)
V_optimal = solution.value[0]
profit_optimal = -objective(solution.value, (params, fresh_feed))
print("Gradient-Based Optimization (optimistix BFGS)")
print("=" * 50)
print(f"Initial V: 1.0 m³")
print(f"Optimal V: {float(V_optimal):.3f} m³")
print(f"Optimal profit: ${float(profit_optimal):.2f}/s")
# Show profit at a few points for comparison
print(f"\nProfit at different volumes:")
for V_test in [0.5, 1.0, 2.0, 3.0, 4.0, 5.0]:
profit = -objective(jnp.array([V_test]), (params, fresh_feed))
print(f" V = {V_test:.1f} m³: ${float(profit):.2f}/s")
Gradient-Based Optimization (optimistix BFGS)
==================================================
Initial V: 1.0 m³
Optimal V: 1.510 m³
Optimal profit: $844.55/s
Profit at different volumes:
V = 0.5 m³: $748.34/s
V = 1.0 m³: $831.97/s
V = 2.0 m³: $838.64/s
V = 3.0 m³: $750.03/s
V = 4.0 m³: $723.08/s
V = 5.0 m³: $687.50/s
# Visualize the profit curve and optimal point
# Note: We limit the range to where the recycle loop converges reliably
V_range = np.linspace(0.5, 2.0, 30)
profits = [-float(objective(jnp.array([V]), (params, fresh_feed))) for V in V_range]
fig, ax = plt.subplots(figsize=(10, 5))
ax.plot(V_range, profits, 'b-', linewidth=2, label='Profit curve')
ax.axvline(x=float(V_optimal), color='r', linestyle='--', alpha=0.7, label=f'Optimal V = {float(V_optimal):.2f} m³')
ax.plot(float(V_optimal), float(profit_optimal), 'ro', markersize=12, label=f'Max profit = ${float(profit_optimal):.0f}/s')
ax.set_xlabel('Reactor Volume (m³)', fontsize=12)
ax.set_ylabel('Profit ($/s)', fontsize=12)
ax.set_title('Gradient-Based Optimization of Reactor Size (BFGS)', fontsize=12)
ax.legend()
ax.grid(True, alpha=0.3)
plt.tight_layout()
Application 3: Uncertainty Propagation#
Question: If rate constant \(k\) is uncertain (±10%), how uncertain is our product flow?
For small uncertainties, linear propagation works:
# Uncertainty in k: ±2% (small enough for linear approximation)
k_mean = float(params['k'])
k_std = 0.02 * k_mean # 2% uncertainty
# Sensitivity of product to k
dFB_dk = float(grads['k'])
# Linear uncertainty propagation
FB_mean = float(F_B_product)
FB_std_linear = abs(dFB_dk) * k_std
print("Uncertainty Propagation (Linear)")
print("=" * 50)
print(f"Rate constant k: {k_mean:.3f} ± {k_std:.5f} (1/s) [±2%]")
print(f"Sensitivity ∂F_B/∂k: {dFB_dk:.4f} mol/s per (1/s)")
print(f"")
print(f"Product F_B: {FB_mean:.4f} ± {FB_std_linear:.4f} mol/s")
print(f"Relative uncertainty: ±{FB_std_linear/FB_mean*100:.2f}%")
Uncertainty Propagation (Linear)
==================================================
Rate constant k: 0.300 ± 0.00600 (1/s) [±2%]
Sensitivity ∂F_B/∂k: 1.9187 mol/s per (1/s)
Product F_B: 9.3864 ± 0.0115 mol/s
Relative uncertainty: ±0.12%
# Verify with Monte Carlo
np.random.seed(42) # For reproducibility
n_samples = 100
k_samples = np.random.normal(k_mean, k_std, n_samples)
FB_samples = []
for k_sample in k_samples:
params_mc = params.copy()
params_mc['k'] = jnp.array(k_sample)
FB_samples.append(float(solve_flowsheet(params_mc, fresh_feed)))
FB_mean_mc = np.mean(FB_samples)
FB_std_mc = np.std(FB_samples)
print(f"Monte Carlo Verification ({n_samples} samples)")
print("=" * 50)
print(f"Product F_B: {FB_mean_mc:.4f} ± {FB_std_mc:.4f} mol/s")
print(f"")
print(f"Comparison:")
print(f" Linear: {FB_mean:.4f} ± {FB_std_linear:.4f} mol/s")
print(f" Monte Carlo: {FB_mean_mc:.4f} ± {FB_std_mc:.4f} mol/s")
print(f"")
print(f" Std dev match: {abs(FB_std_linear - FB_std_mc)/FB_std_linear*100:.1f}% difference")
print(f"")
print(f"Note: Linear propagation is accurate for small uncertainties.")
print(f"For larger uncertainties (>5%), nonlinear effects become significant.")
Monte Carlo Verification (100 samples)
==================================================
Product F_B: 9.3850 ± 0.0105 mol/s
Comparison:
Linear: 9.3864 ± 0.0115 mol/s
Monte Carlo: 9.3850 ± 0.0105 mol/s
Std dev match: 8.9% difference
Note: Linear propagation is accurate for small uncertainties.
For larger uncertainties (>5%), nonlinear effects become significant.
The Magic: Differentiating Through Iteration#
How does JAX differentiate through the recycle iteration?
Implicit Function Theorem: If \(F(x, \theta) = 0\) defines \(x^*(\theta)\), then:
The optimistix library uses this to provide exact gradients through the converged solution, without differentiating through each iteration step.
This is:
Exact (no approximation)
Efficient (doesn’t unroll the iteration)
Stable (works regardless of iteration count)
Where to Go From Here#
You now have the foundation to explore the advanced examples in this repository:
Tutorials#
tutorials/01_jax_fundamentals.ipynb- Deeper dive into JAX autodifftutorials/02_optimization.ipynb- Advanced optimization techniquestutorials/03_differential_equations.ipynb- ODEs in JAX
Examples#
examples/02_cstr_sensitivity.ipynb- More sensitivity analysisexamples/03_optimization.ipynb- Process optimizationexamples/05_technoeconomic_analysis.ipynb- Economic optimizationexamples/06_uncertainty_propagation.ipynb- Advanced UQ methodsexamples/14_parameter_estimation.ipynb- Fitting models to data
Specialized Domains#
examples/bio/- Biomanufacturing processesexamples/04_rare_earth_extraction.ipynb- Separation processes
Summary of the Tutorial Series#
Tutorial |
Key Concept |
|---|---|
00a |
Flowsheets are boxes (units) connected by arrows (streams) |
00b |
Each unit is a set of equations; streams are shared variables |
00c |
DAGs solve forward; recycles need iteration |
00d |
CSTR: \(F_{out} = F_{in} + V \cdot r\) (algebraic) |
00e |
Energy balance couples T and conversion |
00f |
PFR: \(dF/dV = r\) (ODE, needs integration) |
00g |
Flash: Rachford-Rice equation for VLE |
00h |
Sequential solving for connected units |
00i |
Splitters and mixers for parallel/bypass |
00j |
Fixed-point iteration for recycles |
00k |
Differentiability enables optimization, sensitivity, UQ |
Key Takeaways#
Differentiable flowsheets = simulation + automatic gradients
Gradients enable: Optimization, sensitivity analysis, uncertainty propagation
JAX autodiff: Exact gradients, efficient (one backward pass)
Implicit differentiation: Works through converged recycle loops
Foundation for: Parameter estimation, design optimization, real-time control
Congratulations! You’ve completed the introductory tutorial series. You now have a solid mental model of what flowsheets are, how they’re solved, and why making them differentiable opens up powerful new capabilities.