Differentiable Data Structures in JAX#

JAX can differentiate through nested data structures called PyTrees. Understanding how to create custom differentiable data structures is essential for building clean, modular differentiable programs.

Topics covered:

  1. PyTree fundamentals

  2. Why standard Python classes break with JAX

  3. Custom PyTree registration

  4. JAX-compatible dataclasses (flax.struct, equinox)

  5. Static vs dynamic fields

  6. Filtering and partitioning

  7. Chemical engineering application: Stream and Unit classes

import jax
import jax.numpy as jnp
from jax import random, grad, jit, vmap
from jax import tree_util
import equinox as eqx
from typing import NamedTuple, Any, Callable
from dataclasses import dataclass
import matplotlib.pyplot as plt

jax.config.update("jax_enable_x64", True)

print(f"JAX version: {jax.__version__}")
JAX version: 0.8.2

1. PyTree Fundamentals#

A PyTree is a nested structure of containers (dicts, lists, tuples) with array leaves. JAX transformations like grad, jit, and vmap work on PyTrees automatically.

Built-in PyTree nodes:

  • dict

  • list

  • tuple

  • None (empty container)

Leaves (not traversed):

  • jnp.ndarray

  • Python scalars

  • Any non-registered type

# Basic PyTree examples

# Simple dict is a PyTree
params = {
    'W': jnp.array([[1.0, 2.0], [3.0, 4.0]]),
    'b': jnp.array([0.1, 0.2])
}

# Nested dict is also a PyTree
nested_params = {
    'layer1': {'W': jnp.ones((2, 3)), 'b': jnp.zeros(2)},
    'layer2': {'W': jnp.ones((3, 1)), 'b': jnp.zeros(3)}
}

# Inspect PyTree structure
print("Simple params:")
leaves, treedef = tree_util.tree_flatten(params)
print(f"  Leaves: {[l.shape for l in leaves]}")
print(f"  Structure: {treedef}")

print("\nNested params:")
leaves, treedef = tree_util.tree_flatten(nested_params)
print(f"  Leaves: {[l.shape for l in leaves]}")
print(f"  Structure: {treedef}")
WARNING:2026-01-10 20:59:09,072: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.
Simple params:
  Leaves: [(2, 2), (2,)]
  Structure: PyTreeDef({'W': *, 'b': *})

Nested params:
  Leaves: [(2, 3), (2,), (3, 1), (3,)]
  Structure: PyTreeDef({'layer1': {'W': *, 'b': *}, 'layer2': {'W': *, 'b': *}})
# tree_map applies a function to all leaves

# Double all parameters
doubled = jax.tree_util.tree_map(lambda x: 2 * x, params)
print("Original W:\n", params['W'])
print("Doubled W:\n", doubled['W'])

# Get shapes of all leaves
shapes = jax.tree_util.tree_map(lambda x: x.shape, nested_params)
print("\nNested param shapes:", shapes)
Original W:
 [[1. 2.]
 [3. 4.]]
Doubled W:
 [[2. 4.]
 [6. 8.]]

Nested param shapes: {'layer1': {'W': (2, 3), 'b': (2,)}, 'layer2': {'W': (3, 1), 'b': (3,)}}
# Gradients through PyTrees

def loss_fn(params, x):
    """Simple loss using nested params."""
    h = x @ params['layer1']['W'] + params['layer1']['b']
    h = jax.nn.relu(h)
    y = h @ params['layer2']['W'] + params['layer2']['b']
    return jnp.sum(y ** 2)

# Fix: bias dimensions must match layer output dimensions
# layer1: (2, 3) -> output has 3 features, so b needs shape (3,)
# layer2: (3, 1) -> output has 1 feature, so b needs shape (1,)
nested_params = {
    'layer1': {'W': jnp.ones((2, 3)), 'b': jnp.zeros(3)},
    'layer2': {'W': jnp.ones((3, 1)), 'b': jnp.zeros(1)}
}

x = jnp.ones((1, 2))

# Gradient returns PyTree with same structure!
grads = grad(loss_fn)(nested_params, x)

print("Gradient structure matches params:")
print(f"  grads['layer1']['W'].shape = {grads['layer1']['W'].shape}")
print(f"  grads['layer2']['b'].shape = {grads['layer2']['b'].shape}")
Gradient structure matches params:
  grads['layer1']['W'].shape = (2, 3)
  grads['layer2']['b'].shape = (1,)

2. Why Standard Python Classes Break#

By default, JAX treats custom classes as opaque leaves, not as containers to traverse. This causes problems.

# Standard Python class - DOES NOT work with JAX properly

class SimpleLayer:
    """A simple layer - but JAX can't see inside!"""
    def __init__(self, W, b):
        self.W = W
        self.b = b
    
    def __call__(self, x):
        return x @ self.W + self.b

layer = SimpleLayer(
    W=jnp.array([[1.0, 2.0], [3.0, 4.0]]),
    b=jnp.array([0.1, 0.2])
)

# JAX sees this as a single leaf, not a container!
leaves, treedef = tree_util.tree_flatten(layer)
print(f"Leaves: {leaves}")
print(f"JAX sees the entire object as ONE leaf!")
Leaves: [<__main__.SimpleLayer object at 0x70f9a234d250>]
JAX sees the entire object as ONE leaf!
# This causes problems with grad

def loss_with_layer(layer, x):
    y = layer(x)
    return jnp.sum(y ** 2)

x = jnp.ones((1, 2))

# This will fail or give wrong results!
try:
    grads = grad(loss_with_layer)(layer, x)
    print(f"Gradient type: {type(grads)}")
    print("This 'worked' but grads is useless - it's a zero tangent for an opaque object!")
except Exception as e:
    print(f"Error: {e}")
Error: Argument '<__main__.SimpleLayer object at 0x70f9a234d250>' of type <class '__main__.SimpleLayer'> is not a valid JAX type.
# Standard @dataclass has the same problem

@dataclass
class DataclassLayer:
    W: jnp.ndarray
    b: jnp.ndarray
    
    def __call__(self, x):
        return x @ self.W + self.b

dc_layer = DataclassLayer(
    W=jnp.array([[1.0, 2.0], [3.0, 4.0]]),
    b=jnp.array([0.1, 0.2])
)

# Still seen as a single leaf!
leaves, _ = tree_util.tree_flatten(dc_layer)
print(f"Dataclass leaves: {leaves}")
print("Standard @dataclass is also opaque to JAX!")
Dataclass leaves: [DataclassLayer(W=Array([[1., 2.],
       [3., 4.]], dtype=float64), b=Array([0.1, 0.2], dtype=float64))]
Standard @dataclass is also opaque to JAX!

3. Custom PyTree Registration#

We can teach JAX how to handle custom classes by registering them as PyTree nodes.

Two methods:

  1. @jax.tree_util.register_pytree_node_class decorator

  2. jax.tree_util.register_pytree_node function

# Method 1: Decorator

@tree_util.register_pytree_node_class
class RegisteredLayer:
    """A layer properly registered as a PyTree."""
    
    def __init__(self, W, b):
        self.W = W
        self.b = b
    
    def __call__(self, x):
        return x @ self.W + self.b
    
    def tree_flatten(self):
        """
        Returns (children, aux_data).
        children: the array leaves that JAX should traverse
        aux_data: static data needed to reconstruct the object
        """
        children = (self.W, self.b)  # These are the differentiable parts
        aux_data = None  # No static data in this case
        return children, aux_data
    
    @classmethod
    def tree_unflatten(cls, aux_data, children):
        """Reconstruct the object from children and aux_data."""
        W, b = children
        return cls(W, b)

# Now JAX can see inside!
reg_layer = RegisteredLayer(
    W=jnp.array([[1.0, 2.0], [3.0, 4.0]]),
    b=jnp.array([0.1, 0.2])
)

leaves, treedef = tree_util.tree_flatten(reg_layer)
print(f"Leaves: {[l.shape for l in leaves]}")
print(f"TreeDef: {treedef}")
Leaves: [(2, 2), (2,)]
TreeDef: PyTreeDef(CustomNode(RegisteredLayer[None], [*, *]))
# Now gradients work!

def loss_with_registered(layer, x):
    y = layer(x)
    return jnp.sum(y ** 2)

x = jnp.ones((1, 2))
grads = grad(loss_with_registered)(reg_layer, x)

print("Gradients work now!")
print(f"  grad.W.shape = {grads.W.shape}")
print(f"  grad.W = \n{grads.W}")
print(f"  grad.b = {grads.b}")
Gradients work now!
  grad.W.shape = (2, 2)
  grad.W = 
[[ 8.2 12.4]
 [ 8.2 12.4]]
  grad.b = [ 8.2 12.4]
# With auxiliary (static) data

@tree_util.register_pytree_node_class
class LayerWithActivation:
    """Layer with configurable activation (static data)."""
    
    def __init__(self, W, b, activation='relu'):
        self.W = W
        self.b = b
        self.activation = activation  # This is static, not an array
    
    def __call__(self, x):
        y = x @ self.W + self.b
        if self.activation == 'relu':
            return jax.nn.relu(y)
        elif self.activation == 'tanh':
            return jnp.tanh(y)
        else:
            return y
    
    def tree_flatten(self):
        children = (self.W, self.b)
        aux_data = self.activation  # Static data preserved here
        return children, aux_data
    
    @classmethod
    def tree_unflatten(cls, aux_data, children):
        W, b = children
        activation = aux_data
        return cls(W, b, activation)

# Test
layer_relu = LayerWithActivation(
    W=jnp.array([[1.0], [-1.0]]),
    b=jnp.array([0.0]),
    activation='relu'
)

layer_tanh = LayerWithActivation(
    W=jnp.array([[1.0], [-1.0]]),
    b=jnp.array([0.0]),
    activation='tanh'
)

x = jnp.array([[1.0, 1.0]])
print(f"ReLU output: {layer_relu(x)}")
print(f"Tanh output: {layer_tanh(x)}")

# Gradient works and preserves activation
grad_layer = grad(lambda l, x: jnp.sum(l(x)))(layer_relu, x)
print(f"\nGradient preserves activation: {grad_layer.activation}")
ReLU output: [[0.]]
Tanh output: [[0.]]
Gradient preserves activation: relu

4. NamedTuples as PyTrees#

NamedTuple is automatically a PyTree node in JAX - a simple solution for immutable data structures.

# NamedTuple works out of the box!

class LayerParams(NamedTuple):
    W: jnp.ndarray
    b: jnp.ndarray

class NetworkParams(NamedTuple):
    layer1: LayerParams
    layer2: LayerParams

# Create nested structure
params = NetworkParams(
    layer1=LayerParams(W=jnp.ones((2, 3)), b=jnp.zeros(3)),
    layer2=LayerParams(W=jnp.ones((3, 1)), b=jnp.zeros(1))
)

# JAX can see inside!
leaves, treedef = tree_util.tree_flatten(params)
print(f"Leaves: {[l.shape for l in leaves]}")

# Gradient works
def forward(params, x):
    h = jax.nn.relu(x @ params.layer1.W + params.layer1.b)
    return jnp.sum(h @ params.layer2.W + params.layer2.b)

x = jnp.ones((1, 2))
grads = grad(forward)(params, x)

print(f"\nGradient is also a NetworkParams: {type(grads).__name__}")
print(f"grads.layer1.W.shape = {grads.layer1.W.shape}")
Leaves: [(2, 3), (3,), (3, 1), (1,)]

Gradient is also a NetworkParams: NetworkParams
grads.layer1.W.shape = (2, 3)
# NamedTuple for chemical engineering: Stream

class Stream(NamedTuple):
    """A process stream with molar flows, temperature, pressure."""
    F: jnp.ndarray  # Molar flows for each species (mol/s)
    T: float        # Temperature (K)
    P: float        # Pressure (Pa)

# Create a stream
feed = Stream(
    F=jnp.array([10.0, 0.0, 0.0]),  # 10 mol/s of species A
    T=300.0,
    P=101325.0
)

print(f"Feed stream: {feed}")

# Function using stream
def total_flow(stream):
    return jnp.sum(stream.F)

# Gradient w.r.t. stream
grad_stream = grad(total_flow)(feed)
print(f"\nGradient of total flow w.r.t. stream:")
print(f"  dF/dF = {grad_stream.F}")
print(f"  dF/dT = {grad_stream.T}")
Feed stream: Stream(F=Array([10.,  0.,  0.], dtype=float64), T=300.0, P=101325.0)

Gradient of total flow w.r.t. stream:
  dF/dF = [1. 1. 1.]
  dF/dT = 0.0

5. Equinox: The Modern Approach#

Equinox provides eqx.Module, a base class that makes any class a proper PyTree. This is the recommended approach for complex models.

# Equinox Module - clean and simple

class EquinoxLayer(eqx.Module):
    """A layer using Equinox - automatically a PyTree!"""
    W: jnp.ndarray
    b: jnp.ndarray
    
    def __init__(self, in_features, out_features, key):
        self.W = random.normal(key, (in_features, out_features)) * 0.1
        self.b = jnp.zeros(out_features)
    
    def __call__(self, x):
        return x @ self.W + self.b

# Create layer
key = random.PRNGKey(0)
layer = EquinoxLayer(2, 3, key)

# It's a PyTree!
leaves, treedef = tree_util.tree_flatten(layer)
print(f"Leaves: {[l.shape for l in leaves]}")

# Gradients work
def loss(layer, x):
    return jnp.sum(layer(x) ** 2)

x = jnp.ones((1, 2))
grads = grad(loss)(layer, x)
print(f"\nGradient W shape: {grads.W.shape}")
print(f"Gradient b shape: {grads.b.shape}")
Leaves: [(2, 3), (3,)]
Gradient W shape: (2, 3)
Gradient b shape: (3,)
# Static fields with eqx.static_field()

class ConfigurableLayer(eqx.Module):
    """Layer with static configuration."""
    W: jnp.ndarray
    b: jnp.ndarray
    activation: str = eqx.static_field()  # Not a leaf, not differentiated
    use_bias: bool = eqx.static_field()
    
    def __init__(self, in_features, out_features, activation='relu', use_bias=True, *, key):
        self.W = random.normal(key, (in_features, out_features)) * 0.1
        self.b = jnp.zeros(out_features) if use_bias else None
        self.activation = activation
        self.use_bias = use_bias
    
    def __call__(self, x):
        y = x @ self.W
        if self.use_bias and self.b is not None:
            y = y + self.b
        
        if self.activation == 'relu':
            return jax.nn.relu(y)
        elif self.activation == 'tanh':
            return jnp.tanh(y)
        return y

layer = ConfigurableLayer(2, 3, activation='relu', use_bias=True, key=random.PRNGKey(0))

# Only W and b are leaves (activation and use_bias are static)
leaves, _ = tree_util.tree_flatten(layer)
print(f"Number of leaves: {len(leaves)}")
print(f"Leaf shapes: {[l.shape if l is not None else None for l in leaves]}")

# Static fields are preserved through transformations
grads = grad(loss)(layer, x)
print(f"\nGradient preserves static fields:")
print(f"  activation = {grads.activation}")
print(f"  use_bias = {grads.use_bias}")
Number of leaves: 2
Leaf shapes: [(2, 3), (3,)]

Gradient preserves static fields:
  activation = relu
  use_bias = True
/tmp/ipykernel_1711812/3427054361.py:7: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  activation: str = eqx.static_field()  # Not a leaf, not differentiated
/tmp/ipykernel_1711812/3427054361.py:8: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  use_bias: bool = eqx.static_field()

6. Filtering and Partitioning#

Often you need to separate trainable from non-trainable parts, or arrays from non-arrays. Equinox provides powerful filtering utilities.

# eqx.filter - extract only certain leaves

class MixedModel(eqx.Module):
    """Model with trainable and frozen parameters."""
    trainable_W: jnp.ndarray
    trainable_b: jnp.ndarray
    frozen_scale: jnp.ndarray  # We'll treat this as frozen
    name: str = eqx.static_field()
    
    def __call__(self, x):
        return self.frozen_scale * (x @ self.trainable_W + self.trainable_b)

model = MixedModel(
    trainable_W=jnp.ones((2, 3)),
    trainable_b=jnp.zeros(3),
    frozen_scale=jnp.array(2.0),
    name="my_model"
)

# Filter to get only arrays
arrays_only = eqx.filter(model, eqx.is_array)
print("Arrays only:")
print(f"  trainable_W: {arrays_only.trainable_W.shape}")
print(f"  trainable_b: {arrays_only.trainable_b.shape}")
print(f"  frozen_scale: {arrays_only.frozen_scale}")
Arrays only:
  trainable_W: (2, 3)
  trainable_b: (3,)
  frozen_scale: 2.0
/tmp/ipykernel_1711812/2023596013.py:8: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  name: str = eqx.static_field()
# eqx.partition - split into two parts

# Define what's trainable using a filter function
def is_trainable(leaf):
    """Returns True for leaves we want to train."""
    return eqx.is_array(leaf)

# Custom filter: only parameters with 'trainable' in the path
# We'll use a simpler approach: filter by a custom spec

# Create a filter spec that matches the model structure
# Note: static fields must match between model and filter_spec
filter_spec = MixedModel(
    trainable_W=True,   # Train this
    trainable_b=True,   # Train this
    frozen_scale=False, # Don't train this
    name="my_model"     # Must match model's name
)

trainable, frozen = eqx.partition(model, filter_spec)

print("Trainable part:")
print(f"  W: {trainable.trainable_W}")
print(f"  b: {trainable.trainable_b}")
print(f"  frozen_scale: {trainable.frozen_scale}")  # None - filtered out

print("\nFrozen part:")
print(f"  W: {frozen.trainable_W}")  # None - filtered out
print(f"  frozen_scale: {frozen.frozen_scale}")
Trainable part:
  W: [[1. 1. 1.]
 [1. 1. 1.]]
  b: [0. 0. 0.]
  frozen_scale: None

Frozen part:
  W: None
  frozen_scale: 2.0
# Combining partition with gradient computation

def loss_fn(model, x):
    return jnp.sum(model(x) ** 2)

# Only compute gradients for trainable parameters
@eqx.filter_jit
def train_step(model, x, filter_spec):
    # Partition into trainable and frozen
    trainable, frozen = eqx.partition(model, filter_spec)
    
    # Gradient only w.r.t. trainable
    def loss_trainable(trainable_part):
        full_model = eqx.combine(trainable_part, frozen)
        return loss_fn(full_model, x)
    
    loss, grads = eqx.filter_value_and_grad(loss_trainable)(trainable)
    
    return loss, grads

x = jnp.ones((1, 2))
loss, grads = train_step(model, x, filter_spec)

print(f"Loss: {loss:.4f}")
print(f"Gradient for trainable_W: {grads.trainable_W is not None}")
print(f"Gradient for frozen_scale: {grads.frozen_scale}")
Loss: 48.0000
Gradient for trainable_W: True
Gradient for frozen_scale: None
# eqx.filter_grad - simpler approach

# filter_grad only differentiates w.r.t. array leaves by default
@eqx.filter_jit
def simple_train_step(model, x):
    loss, grads = eqx.filter_value_and_grad(loss_fn)(model, x)
    return loss, grads

loss, grads = simple_train_step(model, x)
print(f"Loss: {loss:.4f}")
print(f"All array fields have gradients:")
print(f"  trainable_W gradient: {grads.trainable_W is not None}")
print(f"  frozen_scale gradient: {grads.frozen_scale is not None}")
Loss: 48.0000
All array fields have gradients:
  trainable_W gradient: True
  frozen_scale gradient: True

7. Chemical Engineering Application: Differentiable Units#

Let’s build proper differentiable data structures for process simulation.

# Stream as an Equinox Module

class ProcessStream(eqx.Module):
    """A differentiable process stream."""
    F: jnp.ndarray          # Molar flows [mol/s] for each species
    T: jnp.ndarray          # Temperature [K] (as array for differentiation)
    P: jnp.ndarray          # Pressure [Pa]
    species: tuple = eqx.static_field()  # Species names (static)
    
    def __init__(self, flows_dict, T, P):
        self.species = tuple(flows_dict.keys())
        self.F = jnp.array([flows_dict[s] for s in self.species])
        self.T = jnp.atleast_1d(jnp.asarray(T))
        self.P = jnp.atleast_1d(jnp.asarray(P))
    
    @property
    def total_flow(self):
        return jnp.sum(self.F)
    
    @property
    def mole_fractions(self):
        return self.F / (self.total_flow + 1e-10)
    
    def get_flow(self, species_name):
        idx = self.species.index(species_name)
        return self.F[idx]

# Create a stream
feed = ProcessStream(
    flows_dict={'A': 10.0, 'B': 0.0, 'C': 0.0},
    T=350.0,
    P=101325.0
)

print(f"Feed stream:")
print(f"  Species: {feed.species}")
print(f"  Flows: {feed.F} mol/s")
print(f"  Total: {feed.total_flow} mol/s")
print(f"  T = {feed.T[0]} K, P = {feed.P[0]} Pa")
Feed stream:
  Species: ('A', 'B', 'C')
  Flows: [10.  0.  0.] mol/s
  Total: 10.0 mol/s
  T = 350.0 K, P = 101325.0 Pa
/tmp/ipykernel_1711812/573322110.py:8: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  species: tuple = eqx.static_field()  # Species names (static)
# CSTR as an Equinox Module

class CSTR(eqx.Module):
    """Differentiable CSTR model."""
    V: jnp.ndarray              # Volume [m³]
    k0: jnp.ndarray             # Pre-exponential factor [1/s]
    Ea: jnp.ndarray             # Activation energy [J/mol]
    stoich: jnp.ndarray         # Stoichiometric coefficients
    species: tuple = eqx.static_field()
    R: float = eqx.static_field()  # Gas constant
    
    def __init__(self, V, k0, Ea, stoich_dict, species):
        self.V = jnp.atleast_1d(jnp.asarray(V))
        self.k0 = jnp.atleast_1d(jnp.asarray(k0))
        self.Ea = jnp.atleast_1d(jnp.asarray(Ea))
        self.species = tuple(species)
        self.stoich = jnp.array([stoich_dict.get(s, 0.0) for s in species])
        self.R = 8.314
    
    def reaction_rate(self, C_A, T):
        """Arrhenius reaction rate."""
        k = self.k0 * jnp.exp(-self.Ea / (self.R * T))
        return k * C_A
    
    def __call__(self, inlet: ProcessStream, Q: jnp.ndarray) -> ProcessStream:
        """
        Solve CSTR steady state.
        
        inlet: inlet stream
        Q: volumetric flow rate [m³/s]
        
        Returns: outlet stream
        """
        # Residence time
        tau = self.V / Q
        
        # Inlet concentration of limiting reactant (species A)
        A_idx = self.species.index('A')
        C_A_in = inlet.F[A_idx] / Q
        
        # Solve: C_A = C_A_in / (1 + k*tau)
        k = self.k0 * jnp.exp(-self.Ea / (self.R * inlet.T))
        C_A_out = C_A_in / (1 + k[0] * tau[0])
        
        # Extent of reaction
        extent = (C_A_in - C_A_out) * Q
        
        # Outlet flows
        F_out = inlet.F + self.stoich * extent
        
        # Create outlet stream (same T, P for isothermal)
        outlet = ProcessStream.__new__(ProcessStream)
        object.__setattr__(outlet, 'F', F_out)
        object.__setattr__(outlet, 'T', inlet.T)
        object.__setattr__(outlet, 'P', inlet.P)
        object.__setattr__(outlet, 'species', inlet.species)
        
        return outlet

# Create CSTR
cstr = CSTR(
    V=1.0,
    k0=1e6,
    Ea=50000.0,
    stoich_dict={'A': -1.0, 'B': 1.0, 'C': 0.0},
    species=['A', 'B', 'C']
)

# Solve
Q = jnp.array([0.01])  # m³/s
outlet = cstr(feed, Q)

print(f"CSTR outlet:")
print(f"  F_A = {outlet.F[0]:.4f} mol/s")
print(f"  F_B = {outlet.F[1]:.4f} mol/s")
print(f"  Conversion = {(feed.F[0] - outlet.F[0]) / feed.F[0] * 100:.1f}%")
/tmp/ipykernel_1711812/2363066064.py:9: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  species: tuple = eqx.static_field()
/tmp/ipykernel_1711812/2363066064.py:10: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  R: float = eqx.static_field()  # Gas constant
CSTR outlet:
  F_A = 2.2479 mol/s
  F_B = 7.7521 mol/s
  Conversion = 77.5%
# Differentiate through the CSTR!

def conversion(cstr, feed, Q):
    """Calculate conversion."""
    outlet = cstr(feed, Q)
    X = (feed.F[0] - outlet.F[0]) / feed.F[0]
    return X

# Gradient of conversion w.r.t. CSTR parameters
grad_cstr = grad(conversion)(cstr, feed, Q)

print("Sensitivity of conversion to CSTR parameters:")
print(f"  dX/dV = {grad_cstr.V[0]:.6f} [1/m³]")
print(f"  dX/dk0 = {grad_cstr.k0[0]:.2e} [s]")
print(f"  dX/dEa = {grad_cstr.Ea[0]:.2e} [mol/J]")

# Gradient w.r.t. inlet conditions
grad_feed = grad(conversion, argnums=1)(cstr, feed, Q)

print(f"\nSensitivity to inlet conditions:")
print(f"  dX/dT_in = {grad_feed.T[0]:.6f} [1/K]")
Sensitivity of conversion to CSTR parameters:
  dX/dV = 0.174260 [1/m³]
  dX/dk0 = 1.74e-07 [s]
  dX/dEa = -5.99e-05 [mol/J]
Sensitivity to inlet conditions:
  dX/dT_in = 0.008555 [1/K]
# Optimize CSTR volume to achieve target conversion

def objective(cstr, feed, Q, target_X=0.90):
    """Squared error from target conversion."""
    X = conversion(cstr, feed, Q)
    return (X - target_X) ** 2

# Gradient-based optimization of volume
import optax

# Start with small volume
cstr_opt = CSTR(V=0.1, k0=1e6, Ea=50000.0, 
                stoich_dict={'A': -1.0, 'B': 1.0, 'C': 0.0},
                species=['A', 'B', 'C'])

# Use filter function approach - only optimize V (the first array leaf)
def filter_fn(path, leaf):
    """Only optimize V parameter."""
    path_str = jax.tree_util.keystr(path)
    return '.V' in path_str and eqx.is_array(leaf)

# Create filter spec using jax.tree_util.tree_map_with_path
filter_spec = jax.tree_util.tree_map_with_path(filter_fn, cstr_opt)

optimizer = optax.adam(learning_rate=0.1)
trainable, frozen = eqx.partition(cstr_opt, filter_spec)
opt_state = optimizer.init(trainable)

print("Optimizing CSTR volume for 90% conversion:")
for i in range(100):
    def loss_fn(trainable):
        model = eqx.combine(trainable, frozen)
        return objective(model, feed, Q)
    
    loss, grads = eqx.filter_value_and_grad(loss_fn)(trainable)
    updates, opt_state = optimizer.update(grads, opt_state)
    trainable = eqx.apply_updates(trainable, updates)
    
    # Keep volume positive
    trainable = eqx.tree_at(lambda t: t.V, trainable, jnp.maximum(trainable.V, 0.01))
    
    if i % 20 == 0:
        current_cstr = eqx.combine(trainable, frozen)
        X = conversion(current_cstr, feed, Q)
        print(f"  Iter {i}: V = {trainable.V[0]:.4f} m³, X = {X*100:.1f}%")

final_cstr = eqx.combine(trainable, frozen)
final_X = conversion(final_cstr, feed, Q)
print(f"\nOptimal volume: {final_cstr.V[0]:.4f} m³")
print(f"Achieved conversion: {final_X*100:.2f}%")
Optimizing CSTR volume for 90% conversion:
  Iter 0: V = 0.2000 m³, X = 40.8%
  Iter 20: V = 1.1418 m³, X = 79.7%
  Iter 40: V = 1.3780 m³, X = 82.6%
  Iter 60: V = 1.4801 m³, X = 83.6%
  Iter 80: V = 1.5548 m³, X = 84.3%
Optimal volume: 1.6184 m³
Achieved conversion: 84.81%
# Composing units: CSTR + Separator

class FlashSeparator(eqx.Module):
    """Simple flash separator based on relative volatility."""
    alpha: jnp.ndarray  # Relative volatilities
    split_frac: jnp.ndarray  # Vapor fraction
    species: tuple = eqx.static_field()
    
    def __init__(self, alpha_dict, split_frac, species):
        self.species = tuple(species)
        self.alpha = jnp.array([alpha_dict.get(s, 1.0) for s in species])
        self.split_frac = jnp.atleast_1d(jnp.asarray(split_frac))
    
    def __call__(self, inlet: ProcessStream):
        """Split inlet into vapor and liquid streams."""
        # Simple model: Ki = alpha_i / sum(alpha * x)
        x = inlet.mole_fractions
        K = self.alpha / jnp.sum(self.alpha * x)
        
        # Vapor fraction of each component
        V = self.split_frac[0]
        y_over_x = K  # Simplified
        vapor_frac = V * y_over_x / (1 + V * (y_over_x - 1))
        
        # Split flows
        F_vapor = inlet.F * vapor_frac
        F_liquid = inlet.F - F_vapor
        
        # Create streams
        vapor = ProcessStream.__new__(ProcessStream)
        object.__setattr__(vapor, 'F', F_vapor)
        object.__setattr__(vapor, 'T', inlet.T)
        object.__setattr__(vapor, 'P', inlet.P)
        object.__setattr__(vapor, 'species', inlet.species)
        
        liquid = ProcessStream.__new__(ProcessStream)
        object.__setattr__(liquid, 'F', F_liquid)
        object.__setattr__(liquid, 'T', inlet.T)
        object.__setattr__(liquid, 'P', inlet.P)
        object.__setattr__(liquid, 'species', inlet.species)
        
        return vapor, liquid

# Create separator
flash = FlashSeparator(
    alpha_dict={'A': 3.0, 'B': 1.0, 'C': 0.5},  # A is most volatile
    split_frac=0.5,
    species=['A', 'B', 'C']
)

# Process train: CSTR → Flash
def process_train(cstr, flash, feed, Q):
    """CSTR followed by flash separator."""
    reactor_out = cstr(feed, Q)
    vapor, liquid = flash(reactor_out)
    return vapor, liquid

vapor, liquid = process_train(final_cstr, flash, feed, Q)

print("Process train: CSTR → Flash")
print(f"\nVapor stream:")
for i, s in enumerate(vapor.species):
    print(f"  {s}: {vapor.F[i]:.4f} mol/s")

print(f"\nLiquid stream:")
for i, s in enumerate(liquid.species):
    print(f"  {s}: {liquid.F[i]:.4f} mol/s")
/tmp/ipykernel_1711812/2975803777.py:7: UserWarning: `equinox.static_field` is deprecated in favour of `equinox.field(static=True)`
  species: tuple = eqx.static_field()
Process train: CSTR → Flash

Vapor stream:
  A: 1.0592 mol/s
  B: 3.6809 mol/s
  C: 0.0000 mol/s

Liquid stream:
  A: 0.4603 mol/s
  B: 4.7996 mol/s
  C: 0.0000 mol/s
# Differentiate through the entire process train!

def product_B_in_liquid(cstr, flash, feed, Q):
    """Objective: maximize B in liquid product."""
    _, liquid = process_train(cstr, flash, feed, Q)
    B_idx = liquid.species.index('B')
    return liquid.F[B_idx]

# Gradients w.r.t. both units
grad_cstr, grad_flash = grad(product_B_in_liquid, argnums=(0, 1))(final_cstr, flash, feed, Q)

print("Sensitivity of liquid B flow to process parameters:")
print(f"\nCSTR:")
print(f"  d(F_B)/dV = {grad_cstr.V[0]:.4f}")
print(f"  d(F_B)/dk0 = {grad_cstr.k0[0]:.2e}")

print(f"\nFlash:")
print(f"  d(F_B)/d(alpha) = {grad_flash.alpha}")
print(f"  d(F_B)/d(split_frac) = {grad_flash.split_frac[0]:.4f}")
Sensitivity of liquid B flow to process parameters:

CSTR:
  d(F_B)/dV = 0.1962
  d(F_B)/dk0 = 3.18e-07

Flash:
  d(F_B)/d(alpha) = [ 0.24276931 -0.72830794  0.        ]
  d(F_B)/d(split_frac) = -8.3330

Summary#

Key concepts:

  1. PyTrees are JAX’s fundamental differentiable data structure

    • Built-in: dict, list, tuple

    • Custom classes need registration

  2. Standard Python classes don’t work with JAX transformations

    • JAX sees them as opaque leaves

    • Gradients don’t flow through

  3. Solutions for custom classes:

    • @register_pytree_node_class decorator

    • NamedTuple (works automatically)

    • eqx.Module (recommended)

  4. Static vs dynamic fields:

    • Dynamic: arrays that get differentiated

    • Static: configuration that doesn’t change (eqx.static_field())

  5. Filtering and partitioning:

    • eqx.filter: extract specific leaves

    • eqx.partition/eqx.combine: split and merge

    • eqx.filter_grad: gradient only for arrays

Chemical engineering applications:

  • ProcessStream: differentiable stream with flows, T, P

  • CSTR, FlashSeparator: differentiable unit operations

  • Compose units into differentiable process trains

  • Sensitivity analysis and optimization “for free”

Best practices:

  • Use eqx.Module for complex models

  • Mark configuration as static_field()

  • Use eqx.filter_jit and eqx.filter_grad for clean code

  • Keep arrays as jnp.ndarray (even scalars as 0-d or 1-d arrays)