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:

  1. Represent flowsheets as systems of equations

  2. Solve single units (CSTR, PFR, flash)

  3. Connect units in series and parallel

  4. 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!

\[\frac{\partial \text{(any output)}}{\partial \text{(any input)}}\]

Why Gradients Matter#

Traditional Approach: Finite Differences#

To compute sensitivity of output \(y\) to parameter \(\theta\):

\[\frac{\partial y}{\partial \theta} \approx \frac{y(\theta + \epsilon) - y(\theta)}{\epsilon}\]

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
../_images/2336656399d1bc9ccee87f799573e6b4e382431246ca44af0aef25370a4025f3.png

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!
../_images/2b01cd277a752e993bbba0b5d066bef6ee64250f79e9d2133c78c35683b68652.png

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()
../_images/6941b9775673d27ca770a0a5a696995c51de75f149df5e90fa8afef7c3006ceb.png

Application 3: Uncertainty Propagation#

Question: If rate constant \(k\) is uncertain (±10%), how uncertain is our product flow?

For small uncertainties, linear propagation works:

\[\sigma_y^2 = \left(\frac{\partial y}{\partial k}\right)^2 \sigma_k^2\]
# 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:

\[\frac{dx^*}{d\theta} = -\left(\frac{\partial F}{\partial x}\right)^{-1} \frac{\partial F}{\partial \theta}\]

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 autodiff

  • tutorials/02_optimization.ipynb - Advanced optimization techniques

  • tutorials/03_differential_equations.ipynb - ODEs in JAX

Examples#

  • examples/02_cstr_sensitivity.ipynb - More sensitivity analysis

  • examples/03_optimization.ipynb - Process optimization

  • examples/05_technoeconomic_analysis.ipynb - Economic optimization

  • examples/06_uncertainty_propagation.ipynb - Advanced UQ methods

  • examples/14_parameter_estimation.ipynb - Fitting models to data

Specialized Domains#

  • examples/bio/ - Biomanufacturing processes

  • examples/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#

  1. Differentiable flowsheets = simulation + automatic gradients

  2. Gradients enable: Optimization, sensitivity analysis, uncertainty propagation

  3. JAX autodiff: Exact gradients, efficient (one backward pass)

  4. Implicit differentiation: Works through converged recycle loops

  5. 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.