Dynamic Modeling#
The difflow.dynamic module provides a unified framework for transient (time-dependent) simulation of chemical process units. This enables:
Startup/shutdown analysis: Simulate process dynamics from empty to steady state
Disturbance response: Track how processes respond to feed changes
Control system design: Test controllers in simulation before implementation
Batch process modeling: Simulate time-varying batch operations
Parameter estimation: Fit kinetic parameters to dynamic experimental data
All dynamic simulations are fully differentiable via JAX, enabling gradient-based optimization of dynamic systems.
Table of Contents#
Quick Start#
import jax.numpy as jnp
from difflow.dynamic import integrate, DynamicCSTR, integrate_unit
from difflow.streams import make_stream
# Simple ODE: harmonic oscillator
def oscillator(t, y):
return jnp.array([y[1], -y[0]])
result = integrate(oscillator, jnp.array([1.0, 0.0]), (0.0, 10.0))
print(f"Final position: {result.y_final[0]:.4f}")
# Dynamic reactor simulation
def rate_fn(C, T, params):
return jnp.array([params["k"] * C["A"]])
cstr = DynamicCSTR(
volume=1.0,
rate_fn=rate_fn,
stoich=jnp.array([[-1], [1]]),
species_order=["A", "B"],
rate_params={"k": 0.1},
)
inlet = make_stream({"A": 1.0, "B": 0.0}, T=350.0, P=101325.0)
result = integrate_unit(cstr, {"inlet": inlet}, (0.0, 100.0))
ODE Integration#
The integrate() Function#
The unified interface for ODE integration:
from difflow.dynamic import integrate
def f(t, y):
"""Derivative function: dy/dt = f(t, y)"""
return -0.5 * y # Exponential decay
result = integrate(
f, # Derivative function
y0=jnp.array([1.0]), # Initial state
t_span=(0.0, 10.0), # Time interval
method="RK4", # Integration method
n_steps=100, # Number of steps (fixed-step methods)
)
# Access results
print(f"Final state: {result.y_final}")
print(f"Success: {result.info.success}")
print(f"Steps taken: {result.info.n_steps}")
# Trajectory (intermediate values)
print(f"Time points: {result.trajectory.t.shape}")
print(f"State history: {result.trajectory.y.shape}")
Available Methods#
Method |
Description |
Use Case |
|---|---|---|
|
4th order Runge-Kutta |
General purpose, fixed step |
|
Adaptive RK45 (Dormand-Prince) |
When accuracy matters |
|
Forward Euler |
Simple problems, debugging |
|
Default diffrax solver (Tsit5) |
Advanced features |
|
Dormand-Prince 5(4) |
Good general purpose |
|
Implicit 5th order |
Stiff systems |
Integration Options#
y0 = jnp.array([1.0])
t_span = (0.0, 10.0)
# Fixed-step methods (RK4, Euler)
result = integrate(f, y0, t_span, method="RK4", n_steps=1000)
# Adaptive methods (RK45)
result = integrate(f, y0, t_span, method="RK45", rtol=1e-6, atol=1e-8)
# Diffrax methods
result = integrate(
f, y0, t_span,
method="diffrax:tsit5",
rtol=1e-5,
atol=1e-7,
max_steps=10000,
)
Dynamic Units#
The DynamicUnit Protocol#
All dynamic units implement the DynamicUnit protocol:
from typing import Protocol
class DynamicUnit(Protocol):
@property
def name(self) -> str: ...
@property
def state_spec(self) -> StateSpec: ...
def initial_state(self, inputs: dict) -> Array: ...
def derivatives(self, t: Array, y: Array, inputs: dict) -> Array: ...
def outputs(self, t: Array, y: Array, inputs: dict) -> dict[str, Stream]: ...
DynamicCSTR#
Continuous stirred-tank reactor with reaction kinetics:
from difflow.dynamic import DynamicCSTR
import jax.numpy as jnp
# Define rate law: r = k * exp(-Ea/RT) * C_A
def rate_fn(C, T, params):
k = params["k0"] * jnp.exp(-params["Ea"] / (8.314 * T))
return jnp.array([k * C["A"]]) # Rate of reaction 1
# A -> B (single reaction)
stoich = jnp.array([
[-1], # A consumed
[+1], # B produced
])
cstr = DynamicCSTR(
volume=1.0, # m³
rate_fn=rate_fn,
stoich=stoich,
species_order=["A", "B"],
rate_params={"k0": 1e6, "Ea": 50000.0},
name="reactor",
)
# State variables: [n_A, n_B] (moles of each species)
DynamicTank#
Storage tank with variable holdup:
from difflow.dynamic import DynamicTank
tank = DynamicTank(
max_volume=10.0, # Maximum volume (m³)
species_order=["A", "B"],
name="storage",
)
# State variables: [V, n_A, n_B] (volume + moles; V starts at 1 m³)
Using integrate_unit()#
Convenience wrapper for simulating dynamic units:
from difflow.dynamic import integrate_unit
from difflow.streams import make_stream
inlet = make_stream({"A": 1.0, "B": 0.0}, T=350.0, P=101325.0)
result = integrate_unit(
cstr,
inputs={"inlet": inlet},
t_span=(0.0, 1000.0),
method="RK4",
n_steps=500,
)
# Final moles
n_A_final, n_B_final = result.y_final
print(f"Final A: {n_A_final:.4f} mol")
print(f"Final B: {n_B_final:.4f} mol")
State Specification#
StateVar and StateSpec#
Define state variables with metadata:
from difflow.dynamic import StateVar, StateSpec
# Manual specification
spec = StateSpec([
StateVar("n_A", "moles", "mol", bounds=(0, None)),
StateVar("n_B", "moles", "mol", bounds=(0, None)),
StateVar("T", "temperature", "K", bounds=(200, 600)),
])
print(f"State dimension: {spec.n_states}")
print(f"State names: {spec.names}")
print(f"Index of T: {spec.get_index('T')}")
Factory Functions#
Convenience functions for common state patterns:
from difflow.dynamic import (
molar_states,
concentration_states,
thermal_state,
volume_state,
reactor_states,
)
# Molar holdup for species
spec = molar_states(["A", "B", "C"]) # n_A, n_B, n_C
# Concentration states
spec = concentration_states(["A", "B"]) # C_A, C_B
# Combine specs
spec = molar_states(["A", "B"]) + thermal_state() # n_A, n_B, T
# Complete reactor state
spec = reactor_states(["A", "B"]) # n_A, n_B, T (moles + temperature)
StateVector#
Runtime access to state values by name:
from difflow.dynamic import StateVector
spec = molar_states(["A", "B"]) + thermal_state()
y = jnp.array([1.0, 0.5, 350.0])
state = StateVector(y, spec)
print(f"n_A = {state['n_A']}")
print(f"T = {state['T']}")
# Or use spec directly
idx_A = spec.get_index("n_A")
n_A = y[idx_A]
Dynamic Flowsheets#
Building a Flowsheet#
Connect multiple dynamic units:
from difflow.dynamic import DynamicFlowsheet, DynamicCSTR, DynamicTank
from difflow.streams import make_stream
# Create units
cstr = DynamicCSTR(
volume=1.0,
rate_fn=rate_fn,
stoich=stoich,
species_order=["A", "B"],
rate_params={"k0": 1e6, "Ea": 50000.0},
name="reactor",
)
tank = DynamicTank(
max_volume=10.0,
species_order=["A", "B"],
name="storage",
)
# Build flowsheet
fs = DynamicFlowsheet(species_order=["A", "B"])
# Add feed stream
feed = make_stream({"A": 1.0, "B": 0.0}, T=350.0, P=101325.0)
fs.add_feed("feed", feed)
# Add units with connections
fs.add_unit(cstr, inlet_names=["feed"], outlet_names=["reactor_out"])
fs.add_unit(tank, inlet_names=["reactor_out"], outlet_names=["product"])
# Simulate
result = fs.simulate(t_span=(0.0, 1000.0), method="RK4", n_steps=500)
Accessing Results#
# Combined final state
print(f"Final state: {result.y_final}")
# Per-unit states
reactor_state = result.unit_state_at("reactor")
tank_state = result.unit_state_at("storage")
# Trajectory
print(f"Time points: {result.trajectory.t.shape}")
# Output streams at final time
streams = fs.outputs(result.trajectory.t[-1], result.y_final)
print(f"Product stream: {streams['product']}")
Time-Varying Feeds#
# Define feed as function of time
def feed_schedule(t):
"""Feed rate doubles after t=500."""
base = make_stream({"A": 1.0, "B": 0.0}, T=350.0, P=101325.0)
scale = jnp.where(t > 500, 2.0, 1.0) # traced t: use jnp.where, not `if`
return {k: (v * scale if k.startswith("F_") else v) for k, v in base.items()}
fs.add_feed("feed", feed_schedule) # Pass function instead of stream
Manual Derivative Access#
For custom integration or analysis:
# Get combined initial state
y0 = fs.initial_state()
# Define derivative function
def flowsheet_f(t, y):
return fs.derivatives(t, y)
# Use with any integrator
result = integrate(flowsheet_f, y0, (0.0, 1000.0), method="RK45")
DAE Systems#
Differential-Algebraic Equations combine ODEs with algebraic constraints:
dx/dt = f(t, x, z) # Differential equations
0 = g(t, x, z) # Algebraic constraints
DAE Units#
from difflow.dynamic import DAEUnitBase, AlgebraicSpec, AlgebraicVar
class MyFlashDrum(DAEUnitBase):
@property
def algebraic_spec(self) -> AlgebraicSpec:
return AlgebraicSpec([
AlgebraicVar("V_frac", "vapor_fraction", "-", bounds=(0, 1)),
])
def algebraic_residual(self, t, x, z, inputs):
"""Return g(t,x,z) - should equal zero at solution."""
V_frac = z[0]
# VLE constraint: sum(z_i * (K_i - 1) / (1 + V*(K_i-1))) = 0
residual = self._rachford_rice(x, V_frac)
return jnp.array([residual])
def differential(self, t, x, z, inputs):
"""Return dx/dt given algebraic variables."""
# Material balances using vapor fraction
...
Built-in: DynamicFlashDrum#
from difflow.dynamic import DynamicFlashDrum, integrate_dae
flash = DynamicFlashDrum(
volume=1.0,
species_order=["A", "B"],
K_func=lambda T: jnp.array([2.0, 0.5]), # K-values (A, B) vs. temperature
name="flash",
)
inlet = make_stream({"A": 0.5, "B": 0.5}, T=350.0, P=101325.0)
result = integrate_dae(
flash,
inputs={"inlet": inlet},
t_span=(0.0, 100.0),
method="RK4",
n_steps=200,
)
print(f"Final moles: {result.x_final}")
print(f"Final vapor fraction: {result.z_final}")
Newton Solver#
The algebraic constraints are solved at each time step:
from difflow.dynamic import newton_solve
def residual(z):
"""System of nonlinear equations."""
x, y = z[0], z[1]
return jnp.array([
x**2 + y**2 - 5.0, # Circle
x * y - 2.0, # Hyperbola
])
z0 = jnp.array([2.0, 1.0])
z_solution, info = newton_solve(residual, z0, tol=1e-8, max_iter=50)
print(f"Solution: {z_solution}")
print(f"Converged: {info['converged']}")
print(f"Residual norm: {jnp.linalg.norm(residual(z_solution)):.2e}")
DAE Integration Methods#
# Euler method (simpler, may need smaller steps)
result = integrate_dae(flash, {"inlet": inlet}, (0.0, 100.0), method="Euler", n_steps=1000)
# RK4 method (more accurate)
result = integrate_dae(flash, {"inlet": inlet}, (0.0, 100.0), method="RK4", n_steps=200)
Diffrax Backend#
Diffrax provides advanced ODE/SDE solvers with adaptive step control.
Installation#
pip install diffrax
Basic Usage#
from difflow.dynamic import integrate
# Use diffrax with default solver (Tsit5)
result = integrate(f, y0, t_span, method="diffrax")
# Specify solver
result = integrate(f, y0, t_span, method="diffrax:dopri5")
Available Solvers#
Explicit (non-stiff problems):
dopri5: Dormand-Prince 5(4) - good general purposedopri8: Dormand-Prince 8(7) - higher accuracytsit5: Tsitouras 5(4) - efficient, recommended defaultbosh3: Bogacki-Shampine 3(2)heun: Heun’s method (2nd order, fixed step)euler: Forward Euler (1st order, fixed step)
Implicit (stiff problems):
kvaerno3: 3rd order implicitkvaerno4: 4th order implicitkvaerno5: 5th order implicit - recommended for stiffimplicit_euler: Backward Euler
Stiff Systems#
Chemical kinetics often involve very different time scales (stiff):
from difflow.dynamic import integrate, integrate_stiff
# Robertson problem - classic stiff test
def robertson(t, y):
k1, k2, k3 = 0.04, 3e7, 1e4
A, B, C = y[0], y[1], y[2]
return jnp.array([
-k1*A + k3*B*C,
k1*A - k2*B*B - k3*B*C,
k2*B*B,
])
y0 = jnp.array([1.0, 0.0, 0.0])
# Using implicit solver
result = integrate(
robertson, y0, (0.0, 1e5),
method="diffrax:kvaerno5",
rtol=1e-4, atol=1e-6,
)
# Or use convenience function
result = integrate_stiff(robertson, y0, (0.0, 1e5))
Tolerance Control#
result = integrate(
f, y0, t_span,
method="diffrax:tsit5",
rtol=1e-6, # Relative tolerance
atol=1e-8, # Absolute tolerance
max_steps=10000, # Maximum integration steps
)
Direct Diffrax API#
For more control:
from difflow.dynamic import integrate_diffrax
result = integrate_diffrax(
f, y0, (0.0, 100.0),
solver="tsit5",
rtol=1e-5,
atol=1e-7,
dt0=0.01, # Initial step size
saveat=jnp.linspace(0, 100, 101), # Save at specific times
)
Gradient Computation#
All dynamic simulations are differentiable via JAX.
Gradients Through Integration#
import jax
def loss(y0):
"""Loss based on final state."""
result = integrate(f, y0, (0.0, 10.0), method="RK4")
return jnp.sum(result.y_final**2)
# Gradient w.r.t. initial condition
y0 = jnp.array([1.0, 0.0])
grad_y0 = jax.grad(loss)(y0)
Parameter Optimization#
def simulate_with_params(k):
"""Simulate reactor with rate constant k."""
def rate_fn(C, T, params):
return jnp.array([k * C["A"]])
cstr = DynamicCSTR(
volume=1.0, rate_fn=rate_fn, stoich=stoich,
species_order=["A", "B"], rate_params={},
)
result = integrate_unit(cstr, {"inlet": inlet}, (0.0, 100.0))
return result.y_final[1] # Final product amount
# Optimize for maximum product
grad_k = jax.grad(simulate_with_params)(jnp.array(0.1))
Sensitivity Analysis#
from difflow.dynamic import sensitivity_analysis
def f_p(t, y, params):
"""dy/dt with an explicit parameter array (here the decay rate k)."""
return -params[0] * y
k_nominal = jnp.array([0.5])
result, sens = sensitivity_analysis(
f_p, jnp.array([1.0]), k_nominal, (0.0, 10.0), method="RK4", n_steps=100,
)
# sens = d y_final / d params (Jacobian, shape (n_states, n_params))
Using integrate_with_grad#
Explicit gradient computation:
from difflow.dynamic import integrate_with_grad
# Returns both result and gradient function
result, grad_fn = integrate_with_grad(f, y0, t_span)
# Compute gradient of final state w.r.t. y0
dy_final_dy0 = grad_fn(jnp.ones_like(result.y_final))
API Reference#
Integration Functions#
Function |
Description |
|---|---|
|
Unified ODE integration |
|
Integrate DynamicUnit |
|
Integrate DAE unit |
|
Fixed-step RK4 |
|
Adaptive RK45 |
|
Forward Euler |
Diffrax Functions#
Function |
Description |
|---|---|
|
Direct diffrax integration |
|
Stiff system integration |
|
Dormand-Prince 5(4) |
|
Tsitouras 5(4) |
|
List available solvers |
|
Check if diffrax installed |
Classes#
Class |
Description |
|---|---|
|
Protocol for dynamic units |
|
Base class with utilities |
|
Dynamic CSTR reactor |
|
Storage tank with holdup |
|
Multi-unit flowsheet |
|
Protocol for DAE units |
|
Base class for DAE units |
|
Flash drum with VLE |
State Classes#
Class |
Description |
|---|---|
|
Single state variable spec |
|
Collection of states |
|
Runtime state access |
|
Algebraic variable spec |
|
Collection of algebraic vars |
Result Classes#
Class |
Description |
|---|---|
|
ODE integration result |
|
Time series of states |
|
Solver statistics |
|
DAE integration result |
|
Flowsheet simulation result |