Superstructure Optimization with Differentiable Flowsheets#
Superstructure optimization simultaneously optimizes both the topology (which units to include) and the continuous design variables (temperatures, volumes, etc.) of a chemical process.
Key Idea: Continuous Relaxation#
Instead of discrete unit selection (MINLP), we use continuous relaxation:
Replace binary choices with smooth blending parameters (0-1)
Use sigmoid/softmax for differentiable “soft” selection
Optimize with gradient-based methods (fast!)
Round to discrete topology if needed
This notebook demonstrates three patterns:
Reactor Selection: Choose between CSTR and PFR
Bypass Optimization: Decide whether to skip a unit
Combined Superstructure: Full topology + design optimization
import jax
import jax.numpy as jnp
from jax import Array
import jaxopt
import matplotlib.pyplot as plt
jax.config.update("jax_enable_x64", True)
from difflow.streams import make_stream, get_flows, scale_stream, combine_streams
from difflow.thermo import IdealThermo, SpeciesData
from difflow.units.cstr import CSTR, CSTRParams
from difflow.units.pfr import PFR, PFRParams
from difflow.units.flash import Flash, FlashParams
print("JAX devices:", jax.devices())
WARNING:2026-02-28 15:32:45,201:jax._src.xla_bridge:905: Platform 'mps' is experimental and not all JAX functionality may be correctly supported!
JAX devices: [CpuDevice(id=0)]
Setup: Chemical System#
We’ll use a simple reaction system: A → B (first-order, exothermic)
This models many industrial processes where we need to choose reactor type and operating conditions.
# Define species thermodynamic properties
species_data = {
"A": SpeciesData(
name="A", MW=100.0, Cp_coeffs=(75.0, 0.0, 0.0, 0.0),
Hvap_coeffs=(35000.0, 0.38, 500.0), antoine_coeffs=(10.0, 3000.0, -50.0),
),
"B": SpeciesData(
name="B", MW=100.0, Cp_coeffs=(75.0, 0.0, 0.0, 0.0),
Hvap_coeffs=(30000.0, 0.38, 450.0), antoine_coeffs=(10.0, 2800.0, -40.0),
Hf=-50000.0, # Exothermic reaction
),
}
thermo = IdealThermo(species_data)
species_order = ["A", "B"]
# Stoichiometry: A -> B
stoich = jnp.array([[-1.0], [+1.0]])
# Arrhenius kinetics: r = k * C_A, k = A * exp(-Ea/RT)
# Using moderate kinetics for numerical stability during optimization
def rate_function(C: dict[str, Array], T: Array, params: dict) -> Array:
k = params["A"] * jnp.exp(-params["Ea"] / (8.314 * T))
return jnp.array([k * C["A"]])
# Kinetic parameters chosen for good numerical behavior:
# - At 300 K: k ~ 0.01 s^-1 (slow reaction)
# - At 400 K: k ~ 0.5 s^-1 (moderate reaction)
# - At 500 K: k ~ 10 s^-1 (fast reaction)
rate_params = {"A": jnp.array(1e4), "Ea": jnp.array(30000.0)}
# Standard feed stream
def make_feed():
return make_stream({"A": 10.0, "B": 0.0}, T=300.0, P=101325.0)
print("Chemical system: A -> B (first-order, exothermic)")
print(f"Kinetics: k = {float(rate_params['A']):.1e} * exp(-{float(rate_params['Ea'])/1000:.1f} kJ/mol / RT)")
# Show rate constants at different temperatures
for T in [300.0, 400.0, 500.0]:
k = float(rate_params['A']) * jnp.exp(-float(rate_params['Ea']) / (8.314 * T))
print(f" k({T:.0f} K) = {k:.3f} s^-1")
Chemical system: A -> B (first-order, exothermic)
Kinetics: k = 1.0e+04 * exp(-30.0 kJ/mol / RT)
k(300 K) = 0.060 s^-1
k(400 K) = 1.209 s^-1
k(500 K) = 7.342 s^-1
Pattern 1: Reactor Selection (CSTR vs PFR)#
Problem: Should we use a CSTR or PFR for this reaction?
Approach: Use a continuous selection parameter s ∈ [0, 1]:
s = 0→ pure CSTRs = 1→ pure PFR0 < s < 1→ weighted blend (for optimization)
The sigmoid function provides smooth gradients for optimization.
def create_cstr(V: Array) -> CSTR:
"""Create a CSTR with given volume."""
params = CSTRParams(
V=V,
rate_fn=rate_function,
stoich=stoich,
rate_params=rate_params,
species_order=species_order,
dH_rxn=jnp.array([-50000.0]),
)
return CSTR(params, thermo=thermo, mode="isothermal")
def create_pfr(V: Array) -> PFR:
"""Create a PFR with given volume."""
params = PFRParams(
V=V,
rate_fn=rate_function,
stoich=stoich,
rate_params=rate_params,
species_order=species_order,
dH_rxn=jnp.array([-50000.0]),
)
return PFR(params, thermo=thermo, mode="isothermal")
def blend_streams(stream1: dict, stream2: dict, weight1: Array) -> dict:
"""Blend two streams with given weights."""
weight2 = 1.0 - weight1
blended = {}
for key in stream1:
blended[key] = weight1 * stream1[key] + weight2 * stream2[key]
return blended
def reactor_selection_objective(params: Array) -> Array:
"""
Optimize reactor selection and operating conditions.
params[0]: selection (unconstrained, sigmoid -> [0,1])
params[1]: V_normalized (sigmoid -> [0.01, 2.0] m^3 for numerical stability)
params[2]: T_normalized (sigmoid -> [320, 450] K, avoiding extreme T)
"""
# Decode parameters with bounded transformations for numerical stability
selection = jax.nn.sigmoid(params[0]) # 0 = CSTR, 1 = PFR
# Volume: bounded to [0.01, 2.0] m^3 to avoid extreme values
V = 0.01 + 1.99 * jax.nn.sigmoid(params[1])
# Temperature: bounded to [320, 450] K to stay in reasonable kinetic range
T = 320.0 + 130.0 * jax.nn.sigmoid(params[2])
inlet = make_feed()
# Run both reactors
cstr = create_cstr(V)
pfr = create_pfr(V)
# Volumetric flow for PFR (assume liquid density ~50 mol/m^3)
Q_v = 10.0 / 50.0 # m^3/s
outlet_cstr, info_cstr = cstr(inlet, T_spec=T)
outlet_pfr, info_pfr = pfr(inlet, volumetric_flow=Q_v, T_spec=T)
# Blend outputs based on selection
outlet = blend_streams(outlet_pfr, outlet_cstr, selection)
# Objective: Maximize profit = Revenue - Costs
F_B = outlet["F_B"]
# Economics ($/year basis, 8000 hr/year)
hours_per_year = 8000.0
revenue = 50.0 * F_B * hours_per_year * 3600 / 1e6 # $M/year
capital = 10000.0 * V / 1e6 # $M (annualized)
energy = 100.0 * (T - 300.0) / 1e6 # $M/year
profit = revenue - capital - energy
return -profit # Minimize negative profit
print("Optimizing reactor selection...")
# Use scipy optimizer for more robust optimization
from scipy.optimize import minimize
# Wrapper for scipy
def scipy_objective(x):
return float(reactor_selection_objective(jnp.array(x)))
def scipy_gradient(x):
return jnp.array(jax.grad(reactor_selection_objective)(jnp.array(x)))
# Initial guess: middle of the bounded parameter space
x0 = [0.0, 0.0, 0.0] # V ~ 1.0 m^3, T ~ 385 K
result = minimize(
scipy_objective,
x0,
method='L-BFGS-B',
jac=scipy_gradient,
options={'maxiter': 200, 'disp': False}
)
# Decode solution
selection = float(jax.nn.sigmoid(result.x[0]))
V_opt = 0.01 + 1.99 * float(jax.nn.sigmoid(result.x[1]))
T_opt = 320.0 + 130.0 * float(jax.nn.sigmoid(result.x[2]))
reactor_type = "PFR" if selection > 0.5 else "CSTR"
print(f"\nOptimal topology: {reactor_type} (selection = {selection:.3f})")
print(f" Volume: {V_opt:.3f} m^3")
print(f" Temperature: {T_opt:.1f} K")
print(f" Profit: ${-float(reactor_selection_objective(jnp.array(result.x))):.3f}M/year")
Optimizing reactor selection...
/var/folders/1f/n8r80n0n5zn2l9h7ghmb94fw0000gn/T/ipykernel_36184/4238439050.py:96: DeprecationWarning: scipy.optimize: The `disp` and `iprint` options of the L-BFGS-B solver are deprecated and will be removed in SciPy 1.18.0.
result = minimize(
Optimal topology: PFR (selection = 1.000)
Volume: 1.043 m^3
Temperature: 449.5 K
Profit: $14399.974M/year
# Visualize the selection landscape
# Create a uniform grid in the TRANSFORMED space (what we plot)
# This avoids distortion from the sigmoid nonlinearity
n_grid = 30
# Create uniform grids in the actual plotted coordinates
selection_actual = jnp.linspace(0.01, 0.99, n_grid) # Selection in [0, 1]
V_actual = jnp.linspace(0.1, 1.9, n_grid) # Volume in m^3
sel_plot, V_plot = jnp.meshgrid(selection_actual, V_actual)
# Fixed T at optimal
T_param = result.x[2]
# Convert back to unconstrained space for objective evaluation
# logit is the inverse of sigmoid
def logit(x):
return jnp.log(x / (1.0 - x))
def V_to_param(V):
"""Convert actual volume to unconstrained parameter."""
# V = 0.01 + 1.99 * sigmoid(param)
# (V - 0.01) / 1.99 = sigmoid(param)
normalized = (V - 0.01) / 1.99
return logit(normalized)
def profit_at_actual_point(sel, V):
"""Compute profit given actual selection and volume values."""
sel_param = logit(sel)
V_param = V_to_param(V)
return -reactor_selection_objective(jnp.array([sel_param, V_param, T_param]))
profit_grid = jax.vmap(jax.vmap(profit_at_actual_point))(sel_plot, V_plot)
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
# Left: Profit surface (now with uniform grid in plot coordinates)
ax1 = axes[0]
cs = ax1.contourf(sel_plot, V_plot, profit_grid, levels=20, cmap='RdYlGn')
plt.colorbar(cs, ax=ax1, label='Profit ($M/year)')
ax1.scatter([selection], [V_opt], c='blue', s=200, marker='*',
edgecolors='white', linewidths=2, zorder=10, label=f'Optimum')
ax1.axvline(0.5, color='white', ls='--', lw=2, alpha=0.7, label='CSTR|PFR boundary')
ax1.set_xlabel('Selection (0=CSTR, 1=PFR)')
ax1.set_ylabel('Volume (m^3)')
ax1.set_title('Profit Landscape: Reactor Selection')
ax1.legend(loc='upper right')
# Right: Comparison at optimal V, varying selection
ax2 = axes[1]
sel_fine = jnp.linspace(0.001, 0.999, 100)
profit_vs_sel = jax.vmap(lambda s: profit_at_actual_point(s, V_opt))(sel_fine)
ax2.plot(sel_fine, profit_vs_sel, 'b-', lw=2)
ax2.axvline(0.5, color='gray', ls='--', alpha=0.5)
ax2.scatter([selection], [-float(reactor_selection_objective(jnp.array(result.x)))],
c='red', s=150, marker='*', zorder=10)
ax2.fill_between([0, 0.5], *ax2.get_ylim(), alpha=0.1, color='blue', label='CSTR region')
ax2.fill_between([0.5, 1], *ax2.get_ylim(), alpha=0.1, color='orange', label='PFR region')
ax2.set_xlabel('Selection Parameter')
ax2.set_ylabel('Profit ($M/year)')
ax2.set_title(f'Profit vs Reactor Selection (V={V_opt:.2f} m^3)')
ax2.legend()
ax2.set_xlim(0, 1)
plt.tight_layout()
plt.show()
print(f"\nInterpretation:")
print(f" The optimizer found that {reactor_type} is more profitable.")
print(f" Selection = {selection:.3f} indicates a clear preference.")
Interpretation:
The optimizer found that PFR is more profitable.
Selection = 1.000 indicates a clear preference.
Pattern 2: Bypass Optimization#
Problem: Should some of the feed bypass the reactor?
Bypass can be beneficial when:
Reactor cost is high relative to marginal conversion gain
Downstream separation can handle unreacted feed
Partial conversion is economically optimal
Approach: Optimize a bypass fraction β ∈ [0, 1]
def with_bypass(inlet: dict, unit_fn, bypass_frac: Array, **unit_kwargs):
"""
Route fraction of flow around a unit.
Args:
inlet: Input stream
unit_fn: Callable that processes a stream
bypass_frac: Fraction to bypass (0 = all through unit)
**unit_kwargs: Additional arguments for unit_fn
Returns:
Combined outlet stream, unit info
"""
bypass_frac = jnp.clip(bypass_frac, 0.0, 1.0)
# Split the inlet
bypassed = scale_stream(inlet, bypass_frac)
to_process = scale_stream(inlet, 1.0 - bypass_frac)
# Process the non-bypassed portion
processed, info = unit_fn(to_process, **unit_kwargs)
# Combine bypassed and processed streams
outlet = combine_streams(bypassed, processed)
return outlet, info
def bypass_objective(params: Array) -> Array:
"""
Optimize bypass fraction and reactor conditions.
params[0]: bypass_frac (sigmoid -> [0,1])
params[1]: V_normalized (sigmoid -> [0.01, 2.0] m^3)
params[2]: T_normalized (sigmoid -> [320, 450] K)
"""
bypass_frac = jax.nn.sigmoid(params[0])
V = 0.01 + 1.99 * jax.nn.sigmoid(params[1])
T = 320.0 + 130.0 * jax.nn.sigmoid(params[2])
inlet = make_feed()
cstr = create_cstr(V)
# Apply bypass around CSTR
outlet, info = with_bypass(inlet, cstr, bypass_frac, T_spec=T)
# Economics
F_B = outlet["F_B"]
F_A_unreacted = outlet["F_A"]
hours_per_year = 8000.0
revenue = 50.0 * F_B * hours_per_year * 3600 / 1e6
# Capital scales with (1 - bypass): smaller effective reactor
capital = 10000.0 * V * (1.0 - bypass_frac) / 1e6
energy = 100.0 * (T - 300.0) * (1.0 - bypass_frac) / 1e6
# Penalty for unreacted A (waste disposal cost)
waste_penalty = 5.0 * F_A_unreacted * hours_per_year * 3600 / 1e6
profit = revenue - capital - energy - waste_penalty
return -profit
print("Optimizing bypass fraction...")
# Use scipy optimizer
def scipy_bypass_objective(x):
return float(bypass_objective(jnp.array(x)))
def scipy_bypass_gradient(x):
return jnp.array(jax.grad(bypass_objective)(jnp.array(x)))
x0_bypass = [0.0, 0.0, 0.0]
result_bypass = minimize(
scipy_bypass_objective, x0_bypass, method='L-BFGS-B',
jac=scipy_bypass_gradient, options={'maxiter': 200, 'disp': False}
)
bypass_opt = float(jax.nn.sigmoid(result_bypass.x[0]))
V_bypass = 0.01 + 1.99 * float(jax.nn.sigmoid(result_bypass.x[1]))
T_bypass = 320.0 + 130.0 * float(jax.nn.sigmoid(result_bypass.x[2]))
print(f"\nOptimal bypass: {bypass_opt*100:.1f}%")
print(f" Reactor volume: {V_bypass:.3f} m^3")
print(f" Temperature: {T_bypass:.1f} K")
print(f" Profit: ${-float(bypass_objective(jnp.array(result_bypass.x))):.3f}M/year")
Optimizing bypass fraction...
/var/folders/1f/n8r80n0n5zn2l9h7ghmb94fw0000gn/T/ipykernel_36184/607749390.py:76: DeprecationWarning: scipy.optimize: The `disp` and `iprint` options of the L-BFGS-B solver are deprecated and will be removed in SciPy 1.18.0.
result_bypass = minimize(
Optimal bypass: 0.0%
Reactor volume: 2.000 m^3
Temperature: 450.0 K
Profit: $13933.093M/year
# Compare: with vs without bypass optimization
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
# Left: Profit vs bypass fraction
ax1 = axes[0]
bypass_range = jnp.linspace(-5, 5, 100)
profit_vs_bypass = jax.vmap(
lambda b: -bypass_objective(jnp.array([b, result_bypass.x[1], result_bypass.x[2]]))
)(bypass_range)
ax1.plot(jax.nn.sigmoid(bypass_range) * 100, profit_vs_bypass, 'b-', lw=2)
ax1.scatter([bypass_opt * 100], [-float(bypass_objective(jnp.array(result_bypass.x)))],
c='red', s=150, marker='*', zorder=10, label=f'Optimum: {bypass_opt*100:.1f}%')
ax1.axhline(-float(bypass_objective(jnp.array([-10, result_bypass.x[1], result_bypass.x[2]]))),
color='gray', ls='--', label='No bypass')
ax1.set_xlabel('Bypass Fraction (%)')
ax1.set_ylabel('Profit ($M/year)')
ax1.set_title('Profit vs Bypass Fraction')
ax1.legend()
ax1.grid(True, alpha=0.3)
# Right: Flow diagram visualization
ax2 = axes[1]
ax2.set_xlim(0, 10)
ax2.set_ylim(0, 10)
ax2.set_aspect('equal')
# Draw flow diagram
from matplotlib.patches import Rectangle, FancyArrowPatch, FancyBboxPatch
# Feed
ax2.annotate('Feed\n10 mol/s A', xy=(0.5, 5), fontsize=10, ha='center')
ax2.arrow(1.2, 5, 1.3, 0, head_width=0.2, head_length=0.1, fc='blue', ec='blue')
# Split point
ax2.plot(2.8, 5, 'ko', markersize=10)
# Bypass stream (top)
ax2.annotate(f'{bypass_opt*100:.0f}%', xy=(4, 7.5), fontsize=9, ha='center', color='orange')
ax2.plot([2.8, 2.8, 6, 6], [5, 7, 7, 5.5], 'orange', lw=2)
ax2.arrow(6, 5.7, 0, -0.1, head_width=0.15, head_length=0.1, fc='orange', ec='orange')
# To reactor (bottom)
ax2.arrow(2.8, 5, 0.5, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
# Reactor
reactor = FancyBboxPatch((3.5, 4), 2, 2, boxstyle="round,pad=0.1",
facecolor='lightblue', edgecolor='black', lw=2)
ax2.add_patch(reactor)
ax2.text(4.5, 5, f'CSTR\nV={V_bypass:.2f} m^3\nT={T_bypass:.0f} K',
ha='center', va='center', fontsize=9)
# From reactor
ax2.arrow(5.5, 5, 0.3, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
# Mix point
ax2.plot(6.2, 5, 'ko', markersize=10)
# Product
ax2.arrow(6.2, 5, 1.3, 0, head_width=0.2, head_length=0.1, fc='green', ec='green')
ax2.annotate('Product', xy=(8.5, 5), fontsize=10, ha='center')
ax2.set_title(f'Optimal Flowsheet: {bypass_opt*100:.0f}% Bypass')
ax2.axis('off')
plt.tight_layout()
plt.show()
Pattern 3: Combined Superstructure#
Problem: Optimize the full process topology:
Choose reactor type (CSTR or PFR)
Decide bypass fraction
Include downstream flash separator
Optimize all continuous variables
This represents a more realistic superstructure optimization problem.
def full_superstructure_objective(params: Array) -> Array:
"""
Full superstructure: Feed -> [Bypass?] -> [CSTR or PFR] -> Flash -> Products
Decision variables:
- params[0]: reactor_selection (0=CSTR, 1=PFR)
- params[1]: bypass_frac
- params[2]: V_reactor (sigmoid -> [0.01, 2.0] m^3)
- params[3]: T_reactor (sigmoid -> [320, 450] K)
- params[4]: T_flash (sigmoid -> [300, 400] K)
"""
# Decode parameters with bounded transformations
reactor_sel = jax.nn.sigmoid(params[0])
bypass_frac = jax.nn.sigmoid(params[1])
V = 0.01 + 1.99 * jax.nn.sigmoid(params[2])
T_reactor = 320.0 + 130.0 * jax.nn.sigmoid(params[3])
T_flash = 300.0 + 100.0 * jax.nn.sigmoid(params[4]) # 300-400 K
inlet = make_feed()
Q_v = 10.0 / 50.0 # Volumetric flow for PFR
# === Bypass split ===
bypassed = scale_stream(inlet, bypass_frac)
to_reactor = scale_stream(inlet, 1.0 - bypass_frac)
# === Reactor selection ===
cstr = create_cstr(V)
pfr = create_pfr(V)
out_cstr, _ = cstr(to_reactor, T_spec=T_reactor)
out_pfr, _ = pfr(to_reactor, volumetric_flow=Q_v * (1.0 - bypass_frac + 0.01), T_spec=T_reactor)
# Blend based on selection
reactor_out = blend_streams(out_pfr, out_cstr, reactor_sel)
# === Combine with bypass ===
combined = combine_streams(bypassed, reactor_out)
# === Flash separator ===
flash_params = FlashParams(species_order=species_order)
flash = Flash(flash_params, thermo=thermo)
liquid, vapor, flash_info = flash(combined, T=T_flash)
# === Economics ===
# Product B is in the liquid (higher boiling point = lower vapor pressure)
F_B_product = liquid["F_B"]
F_A_waste = vapor["F_A"] # Unreacted A goes to vapor (waste)
hours_per_year = 8000.0
seconds_per_year = hours_per_year * 3600
revenue = 50.0 * F_B_product * seconds_per_year / 1e6
# Capital costs
reactor_capital = 10000.0 * V * (1.0 - bypass_frac) / 1e6
flash_capital = 5000.0 / 1e6 # Fixed cost for flash
# Operating costs
reactor_energy = 100.0 * (T_reactor - 300.0) * (1.0 - bypass_frac) / 1e6
flash_energy = 50.0 * jnp.abs(T_flash - combined["T"]) / 1e6
waste_disposal = 10.0 * F_A_waste * seconds_per_year / 1e6
profit = revenue - reactor_capital - flash_capital - reactor_energy - flash_energy - waste_disposal
return -profit
print("Optimizing full superstructure...")
print("Variables: reactor_type, bypass_frac, V, T_reactor, T_flash\n")
# Use scipy optimizer for robustness
def scipy_full_objective(x):
return float(full_superstructure_objective(jnp.array(x)))
def scipy_full_gradient(x):
return jnp.array(jax.grad(full_superstructure_objective)(jnp.array(x)))
x0_full = [0.0, -2.0, 0.0, 0.0, 0.0] # Low bypass, moderate V and T
result_full = minimize(
scipy_full_objective, x0_full, method='L-BFGS-B',
jac=scipy_full_gradient, options={'maxiter': 500, 'disp': False}
)
# Decode solution
reactor_sel_opt = float(jax.nn.sigmoid(result_full.x[0]))
bypass_opt_full = float(jax.nn.sigmoid(result_full.x[1]))
V_opt_full = 0.01 + 1.99 * float(jax.nn.sigmoid(result_full.x[2]))
T_reactor_opt = 320.0 + 130.0 * float(jax.nn.sigmoid(result_full.x[3]))
T_flash_opt = 300.0 + 100.0 * float(jax.nn.sigmoid(result_full.x[4]))
reactor_type_opt = "PFR" if reactor_sel_opt > 0.5 else "CSTR"
print("="*50)
print("OPTIMAL SUPERSTRUCTURE")
print("="*50)
print(f"\nTopology decisions:")
print(f" Reactor type: {reactor_type_opt} (selection = {reactor_sel_opt:.3f})")
print(f" Bypass: {bypass_opt_full*100:.1f}%")
print(f"\nDesign variables:")
print(f" Reactor volume: {V_opt_full:.3f} m^3")
print(f" Reactor temperature: {T_reactor_opt:.1f} K")
print(f" Flash temperature: {T_flash_opt:.1f} K")
print(f"\nEconomics:")
print(f" Profit: ${-float(full_superstructure_objective(jnp.array(result_full.x))):.3f}M/year")
Optimizing full superstructure...
Variables: reactor_type, bypass_frac, V, T_reactor, T_flash
/var/folders/1f/n8r80n0n5zn2l9h7ghmb94fw0000gn/T/ipykernel_36184/2734987612.py:81: DeprecationWarning: scipy.optimize: The `disp` and `iprint` options of the L-BFGS-B solver are deprecated and will be removed in SciPy 1.18.0.
result_full = minimize(
==================================================
OPTIMAL SUPERSTRUCTURE
==================================================
Topology decisions:
Reactor type: PFR (selection = 1.000)
Bypass: 0.0%
Design variables:
Reactor volume: 1.041 m^3
Reactor temperature: 449.8 K
Flash temperature: 399.9 K
Economics:
Profit: $14399.966M/year
# Sensitivity analysis: How does profit change with each decision?
fig, axes = plt.subplots(2, 2, figsize=(12, 10))
# 1. Reactor selection sensitivity
ax1 = axes[0, 0]
sel_range = jnp.linspace(-5, 5, 50)
profit_vs_sel = jax.vmap(
lambda s: -full_superstructure_objective(
jnp.array(result_full.x).at[0].set(s)
)
)(sel_range)
ax1.plot(jax.nn.sigmoid(sel_range), profit_vs_sel, 'b-', lw=2)
ax1.axvline(0.5, color='gray', ls='--', alpha=0.5)
ax1.scatter([reactor_sel_opt], [-float(full_superstructure_objective(jnp.array(result_full.x)))],
c='red', s=100, marker='*', zorder=10)
ax1.set_xlabel('Reactor Selection (0=CSTR, 1=PFR)')
ax1.set_ylabel('Profit ($M/year)')
ax1.set_title('Sensitivity: Reactor Type')
ax1.grid(True, alpha=0.3)
# 2. Bypass sensitivity
ax2 = axes[0, 1]
bypass_range = jnp.linspace(-5, 5, 50)
profit_vs_bypass = jax.vmap(
lambda b: -full_superstructure_objective(
jnp.array(result_full.x).at[1].set(b)
)
)(bypass_range)
ax2.plot(jax.nn.sigmoid(bypass_range) * 100, profit_vs_bypass, 'g-', lw=2)
ax2.scatter([bypass_opt_full * 100], [-float(full_superstructure_objective(jnp.array(result_full.x)))],
c='red', s=100, marker='*', zorder=10)
ax2.set_xlabel('Bypass Fraction (%)')
ax2.set_ylabel('Profit ($M/year)')
ax2.set_title('Sensitivity: Bypass Fraction')
ax2.grid(True, alpha=0.3)
# 3. Volume sensitivity
ax3 = axes[1, 0]
V_range = jnp.linspace(-3, 3, 50)
profit_vs_V = jax.vmap(
lambda v: -full_superstructure_objective(
jnp.array(result_full.x).at[2].set(v)
)
)(V_range)
V_actual = 0.01 + 1.99 * jax.nn.sigmoid(V_range)
ax3.plot(V_actual, profit_vs_V, 'm-', lw=2)
ax3.scatter([V_opt_full], [-float(full_superstructure_objective(jnp.array(result_full.x)))],
c='red', s=100, marker='*', zorder=10)
ax3.set_xlabel('Reactor Volume (m^3)')
ax3.set_ylabel('Profit ($M/year)')
ax3.set_title('Sensitivity: Reactor Volume')
ax3.grid(True, alpha=0.3)
# 4. Temperature sensitivity
ax4 = axes[1, 1]
T_range = jnp.linspace(-5, 5, 50)
profit_vs_T = jax.vmap(
lambda t: -full_superstructure_objective(
jnp.array(result_full.x).at[3].set(t)
)
)(T_range)
T_actual = 320 + 130 * jax.nn.sigmoid(T_range)
ax4.plot(T_actual, profit_vs_T, 'orange', lw=2)
ax4.scatter([T_reactor_opt], [-float(full_superstructure_objective(jnp.array(result_full.x)))],
c='red', s=100, marker='*', zorder=10)
ax4.set_xlabel('Reactor Temperature (K)')
ax4.set_ylabel('Profit ($M/year)')
ax4.set_title('Sensitivity: Reactor Temperature')
ax4.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
# Final visualization: Optimal flowsheet diagram
fig, ax = plt.subplots(figsize=(14, 6))
ax.set_xlim(0, 14)
ax.set_ylim(0, 8)
# Title
ax.text(7, 7.5, 'Optimal Superstructure Configuration',
ha='center', fontsize=14, fontweight='bold')
# Feed
ax.annotate('FEED\n10 mol/s A\n300 K', xy=(0.8, 4), fontsize=9, ha='center',
bbox=dict(boxstyle='round', facecolor='lightblue', alpha=0.8))
ax.arrow(1.6, 4, 0.8, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
# Split point
ax.plot(2.7, 4, 'ko', markersize=12)
ax.text(2.7, 3.3, 'Split', ha='center', fontsize=8)
# Bypass stream
bypass_pct = bypass_opt_full * 100
ax.plot([2.7, 2.7, 8.5, 8.5], [4, 6, 6, 4.3], 'orange', lw=2)
ax.annotate(f'Bypass\n{bypass_pct:.0f}%', xy=(5.5, 6.3), fontsize=9, ha='center', color='darkorange')
# To reactor
ax.arrow(2.7, 4, 0.8, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
ax.annotate(f'{100-bypass_pct:.0f}%', xy=(3.2, 3.5), fontsize=8, ha='center')
# Reactor
reactor_color = 'lightgreen' if reactor_type_opt == 'PFR' else 'lightyellow'
reactor = plt.Rectangle((3.8, 2.8), 2.4, 2.4, facecolor=reactor_color,
edgecolor='black', lw=2)
ax.add_patch(reactor)
ax.text(5, 4, f'{reactor_type_opt}\nV = {V_opt_full:.2f} m^3\nT = {T_reactor_opt:.0f} K',
ha='center', va='center', fontsize=9)
# From reactor
ax.arrow(6.2, 4, 0.8, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
# Mix point
ax.plot(8.5, 4, 'ko', markersize=12)
ax.text(8.5, 3.3, 'Mix', ha='center', fontsize=8)
# To flash
ax.arrow(8.5, 4, 0.8, 0, head_width=0.15, head_length=0.1, fc='blue', ec='blue')
# Flash separator
flash = plt.Rectangle((9.5, 2.5), 1.8, 3, facecolor='lightcyan',
edgecolor='black', lw=2)
ax.add_patch(flash)
ax.text(10.4, 4, f'FLASH\nT = {T_flash_opt:.0f} K', ha='center', va='center', fontsize=9)
# Vapor out (top)
ax.arrow(10.4, 5.5, 0, 0.8, head_width=0.15, head_length=0.1, fc='red', ec='red')
ax.annotate('Vapor\n(waste A)', xy=(10.4, 7), fontsize=9, ha='center', color='red')
# Liquid out (bottom)
ax.arrow(10.4, 2.5, 0, -0.8, head_width=0.15, head_length=0.1, fc='green', ec='green')
ax.annotate('Liquid\n(product B)', xy=(10.4, 1), fontsize=9, ha='center', color='darkgreen')
# Economics summary
profit = -float(full_superstructure_objective(jnp.array(result_full.x)))
ax.text(13, 4, f'PROFIT\n${profit:.2f}M/yr', ha='center', va='center', fontsize=11,
bbox=dict(boxstyle='round', facecolor='lightgreen', alpha=0.9, edgecolor='green', lw=2))
ax.axis('off')
plt.tight_layout()
plt.show()
Summary#
This notebook demonstrated superstructure optimization in difflow using continuous relaxation:
Pattern |
Decision |
Implementation |
|---|---|---|
Unit Selection |
CSTR vs PFR |
Sigmoid-weighted blend of outputs |
Bypass |
Include/skip unit |
|
Multi-way |
Choose among N options |
Softmax-weighted blend |
Key Techniques#
Parameter transformations: Use sigmoid/exp to constrain parameters
Smooth blending: Run all alternatives, blend outputs
Gradient-based optimization: JAX autodiff + scipy.optimize
Post-processing: Round to discrete topology if needed
Advantages#
Fast: Continuous optimization (100-1000x faster than MINLP)
Differentiable: Exact gradients, no finite differences
Flexible: Combine any units, add constraints via penalties
Scalable: JIT compilation, GPU-ready
Extensions#
Add more unit alternatives (distillation, extraction, etc.)
Include recycle streams via
FlowsheetclassMulti-objective optimization (Pareto front)
Uncertainty-aware optimization
Troubleshooting Guide#
Superstructure optimization with differentiable flowsheets can encounter numerical issues. Here are common problems and solutions:
1. “Maximum number of solver steps reached”#
Symptom: Error like EquinoxRuntimeError: The maximum number of solver steps was reached
Cause: The internal ODE/DAE solvers (diffrax) or root-finding solvers (optimistix) exhaust their step limit when exploring difficult parameter regions.
Solutions:
Bound your parameters: Use sigmoid transformations to keep parameters in numerically stable ranges
Choose moderate kinetics: Very fast reactions (high Ea, high A) create stiff equations
Avoid extreme volumes: Very large or very small reactor volumes can cause issues
# BAD: Unbounded log-space
V = jnp.exp(params[1]) # Can grow to infinity!
# GOOD: Bounded with sigmoid
V = 0.01 + 1.99 * jax.nn.sigmoid(params[1]) # Always in [0.01, 2.0]
2. “Linear solver returned non-finite (NaN or inf)”#
Symptom: NaN/Inf errors from lineax during Newton iterations
Cause: The Jacobian matrix becomes singular or ill-conditioned, typically when:
Reaction rates approach zero (low T, depleted reactants)
Reaction rates approach infinity (high T)
Near-complete conversion (outlet flows approach zero)
Solutions:
Bound temperature range: Avoid extreme temperatures where kinetics are too fast/slow
Add small regularization: Ensure flows never exactly reach zero
Use moderate kinetics: Pre-exponential factor of 1e4-1e5 and Ea of 20-40 kJ/mol work well
# Good kinetics for optimization
rate_params = {"A": jnp.array(1e4), "Ea": jnp.array(30000.0)}
# Temperature bounded to reasonable range
T = 320.0 + 130.0 * jax.nn.sigmoid(params[2]) # [320, 450] K
3. Optimizer doesn’t converge or finds wrong solution#
Symptom: scipy.optimize.minimize returns with success=False or finds a local minimum
Solutions:
Try multiple starting points: The objective may have local minima
Check gradients: Verify gradients are finite using
jax.grad(objective)(x0)Simplify first: Test individual units before the full superstructure
Use robust optimizer:
method='L-BFGS-B'with explicitjac=works well
# Test objective at initial point first
x0 = jnp.array([0.0, 0.0, 0.0])
print(f"Objective at x0: {objective(x0)}")
print(f"Gradient at x0: {jax.grad(objective)(x0)}")
4. PFR fails with zero volumetric flow#
Symptom: Error when bypass fraction is high and volumetric flow approaches zero
Cause: The PFR requires positive volumetric flow for integration
Solution: Add a small offset to ensure non-zero flow
# Ensure minimum volumetric flow to PFR
Q_v_pfr = Q_v * (1.0 - bypass_frac + 0.01) # Always > 0
5. Very large profit values#
Symptom: Optimization returns unrealistic profit (e.g., $14000M/year)
Cause: Economic parameters may not be scaled realistically
Solution: Check units and scaling of revenue/cost terms, or focus on relative comparisons rather than absolute values
Best Practices Summary#
Issue |
Solution |
|---|---|
Unbounded parameters |
Use sigmoid to bound to physical range |
Stiff kinetics |
Moderate A (1e4) and Ea (30 kJ/mol) |
Zero flows |
Add small offsets (1e-3 to 0.01) |
Solver failures |
Test components individually first |
Local minima |
Try multiple starting points |
Slow convergence |
Use scipy L-BFGS-B with explicit gradient |