Tutorial 22: Distributed and Parallel Differentiation#
This tutorial covers JAX’s tools for parallel and distributed automatic differentiation across multiple devices (CPUs, GPUs, TPUs). We’ll explore:
Understanding JAX devices and parallelism
pmap- Parallel map for data parallelismGradients with
pmap- distributed trainingjax.sharding- Modern array sharding APIshard_map- Explicit SPMD programmingGradient accumulation strategies
Chemical engineering example: Parallel experiment optimization
import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, pmap
from jax.sharding import NamedSharding, Mesh, PartitionSpec as P, SingleDeviceSharding
from jax.experimental import mesh_utils
from jax.experimental.shard_map import shard_map
import optax
from functools import partial
# Check available devices
print(f"Available devices: {jax.devices()}")
print(f"Number of devices: {jax.device_count()}")
print(f"Local devices: {jax.local_devices()}")
/tmp/ipykernel_1713579/391649778.py:6: DeprecationWarning: jax.experimental.shard_map is deprecated in v0.8.0. Used jax.shard_map instead.
from jax.experimental.shard_map import shard_map
WARNING:2026-01-10 21:00:29,246: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.
Available devices: [CpuDevice(id=0)]
Number of devices: 1
Local devices: [CpuDevice(id=0)]
1. Understanding JAX Devices and Parallelism#
JAX can distribute computation across multiple devices:
Data parallelism: Same model, different data batches on each device
Model parallelism: Different parts of the model on different devices
Pipeline parallelism: Different stages on different devices
Key concepts:
Device: A single accelerator (CPU, GPU, TPU core)
Mesh: A logical arrangement of devices
Sharding: How arrays are distributed across devices
# Create a simple array and examine its sharding
x = jnp.arange(16).reshape(4, 4)
print(f"Array shape: {x.shape}")
print(f"Array devices: {x.devices()}")
print(f"Array sharding: {x.sharding}")
# By default, arrays are replicated on a single device
# We can explicitly control placement
Array shape: (4, 4)
Array devices: {CpuDevice(id=0)}
Array sharding: SingleDeviceSharding(device=CpuDevice(id=0), memory_kind=device)
# For this tutorial, we'll simulate multiple devices using CPU
# In practice, you'd have multiple GPUs/TPUs
# Get device count (may be 1 on a single machine)
n_devices = jax.local_device_count()
print(f"Local device count: {n_devices}")
# Note: On a single CPU machine, pmap will still work but won't provide
# actual parallelism. The concepts translate directly to multi-GPU/TPU.
Local device count: 1
2. pmap - Parallel Map for Data Parallelism#
pmap (parallel map) executes a function on multiple devices simultaneously. Each device processes a different slice of the input data.
# pmap applies function across first axis, one element per device
parallel_fn = pmap(fn)
results = parallel_fn(batched_inputs) # shape[0] == n_devices
# Simple pmap example
def square(x):
return x ** 2
# Create data with leading dimension matching device count
n_devices = jax.local_device_count()
data = jnp.arange(n_devices * 4).reshape(n_devices, 4)
print(f"Input shape: {data.shape}")
print(f"Input data:\n{data}")
# pmap the function
parallel_square = pmap(square)
result = parallel_square(data)
print(f"\nOutput shape: {result.shape}")
print(f"Output data:\n{result}")
Input shape: (1, 4)
Input data:
[[0 1 2 3]]
Output shape: (1, 4)
Output data:
[[0 1 4 9]]
# pmap with multiple arguments
def weighted_sum(x, w):
return jnp.sum(x * w)
x_batched = jnp.ones((n_devices, 10))
w_batched = jnp.arange(n_devices * 10).reshape(n_devices, 10) / 10
parallel_weighted_sum = pmap(weighted_sum)
results = parallel_weighted_sum(x_batched, w_batched)
print(f"Weighted sums per device: {results}")
Weighted sums per device: [4.5]
# Combining pmap with vmap for nested parallelism
# pmap across devices, vmap across batch within each device
def process_sample(x):
"""Process a single sample."""
return jnp.sum(x ** 2)
# vmap for batch processing, pmap for device parallelism
batched_process = vmap(process_sample) # Vectorize over samples
parallel_batched = pmap(batched_process) # Parallelize over devices
# Data: (n_devices, batch_per_device, features)
batch_per_device = 8
features = 16
data = jax.random.normal(jax.random.PRNGKey(0), (n_devices, batch_per_device, features))
results = parallel_batched(data)
print(f"Input shape: {data.shape}")
print(f"Output shape: {results.shape}")
print(f"Total samples processed: {n_devices * batch_per_device}")
Input shape: (1, 8, 16)
Output shape: (1, 8)
Total samples processed: 8
3. Gradients with pmap - Distributed Training#
The key to distributed gradient computation is:
Each device computes gradients on its data shard
Gradients are synchronized across devices (all-reduce)
All devices apply the same update
JAX provides jax.lax.pmean, psum, etc. for collective operations.
# Distributed gradient computation example
def model(params, x):
"""Simple linear model: y = Wx + b."""
return x @ params['W'] + params['b']
def loss_fn(params, x, y):
"""MSE loss."""
pred = model(params, x)
return jnp.mean((pred - y) ** 2)
def compute_gradients(params, x, y):
"""Compute gradients for a batch."""
return grad(loss_fn)(params, x, y)
# Distributed training step
@partial(pmap, axis_name='devices')
def distributed_train_step(params, x, y):
"""
Distributed training step:
1. Compute gradients on local data
2. Average gradients across devices (all-reduce)
3. Return synchronized gradients
"""
grads = compute_gradients(params, x, y)
# Average gradients across all devices
# pmean computes mean across the named axis
grads = jax.lax.pmean(grads, axis_name='devices')
return grads
# Initialize parameters (replicated across devices)
key = jax.random.PRNGKey(0)
params = {
'W': jax.random.normal(key, (10, 5)),
'b': jnp.zeros(5)
}
# Replicate params across devices
replicated_params = jax.tree.map(lambda x: jnp.stack([x] * n_devices), params)
print(f"Replicated W shape: {replicated_params['W'].shape}")
# Create sharded data
batch_size = n_devices * 32
x_data = jax.random.normal(jax.random.PRNGKey(1), (batch_size, 10))
y_data = jax.random.normal(jax.random.PRNGKey(2), (batch_size, 5))
# Reshape for pmap: (n_devices, batch_per_device, ...)
x_sharded = x_data.reshape(n_devices, -1, 10)
y_sharded = y_data.reshape(n_devices, -1, 5)
# Compute distributed gradients
grads = distributed_train_step(replicated_params, x_sharded, y_sharded)
print(f"\nGradient W shape: {grads['W'].shape}")
print(f"Gradients are identical across devices: {jnp.allclose(grads['W'][0], grads['W'][-1])}")
Replicated W shape: (1, 10, 5)
Gradient W shape: (1, 10, 5)
Gradients are identical across devices: True
# Full distributed training loop
@partial(pmap, axis_name='devices')
def distributed_update(params, x, y, learning_rate):
"""Complete distributed update step."""
# Compute and synchronize gradients
grads = compute_gradients(params, x, y)
grads = jax.lax.pmean(grads, axis_name='devices')
# Apply gradient update (same on all devices)
new_params = jax.tree.map(
lambda p, g: p - learning_rate * g,
params, grads
)
# Compute loss for monitoring
loss = loss_fn(params, x, y)
avg_loss = jax.lax.pmean(loss, axis_name='devices')
return new_params, avg_loss
# Training loop
params = replicated_params
learning_rate = jnp.array([0.01] * n_devices) # Must be replicated
print("Distributed Training:")
for epoch in range(5):
params, loss = distributed_update(params, x_sharded, y_sharded, learning_rate)
print(f"Epoch {epoch + 1}: Loss = {loss[0]:.4f}")
Distributed Training:
Epoch 1: Loss = 12.0208
Epoch 2: Loss = 11.8890
Epoch 3: Loss = 11.7592
Epoch 4: Loss = 11.6313
Epoch 5: Loss = 11.5052
# Collective operations available in pmap
@partial(pmap, axis_name='i')
def collective_ops_demo(x):
"""Demonstrate collective operations."""
return {
'psum': jax.lax.psum(x, axis_name='i'), # Sum across devices
'pmean': jax.lax.pmean(x, axis_name='i'), # Mean across devices
'pmax': jax.lax.pmax(x, axis_name='i'), # Max across devices
'pmin': jax.lax.pmin(x, axis_name='i'), # Min across devices
}
# Each device has different values
values = jnp.arange(n_devices, dtype=jnp.float32) + 1
results = collective_ops_demo(values)
print("Collective Operations Demo:")
print(f"Input values: {values}")
for op, result in results.items():
print(f"{op}: {result}")
Collective Operations Demo:
Input values: [1.]
pmax: [1.]
pmean: [1.]
pmin: [1.]
psum: [1.]
4. jax.sharding - Modern Array Sharding API#
JAX’s sharding API provides flexible control over array distribution:
NamedSharding: Shard using named mesh axes (recommended)SingleDeviceSharding: Place array on a single devicejax.device_putwith sharding: Explicitly place arrays
# NamedSharding - distribute array across devices using mesh
devices = jax.devices()
n_devices = len(devices)
if n_devices >= 1:
# Create a mesh and use NamedSharding
mesh = Mesh(devices, axis_names=('data',))
sharding = NamedSharding(mesh, P('data', None))
# Create and shard an array
arr = jnp.arange(n_devices * 4).reshape(n_devices, 4)
sharded_arr = jax.device_put(arr, sharding)
print(f"Original array shape: {arr.shape}")
print(f"Sharding: {sharded_arr.sharding}")
print(f"Array:\n{sharded_arr}")
Original array shape: (1, 4)
Sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec('data', None), memory_kind=device)
Array:
[[0 1 2 3]]
# NamedSharding with Mesh - more expressive sharding
# Create a 1D mesh of devices
mesh = Mesh(jax.devices(), axis_names=('data',))
# PartitionSpec specifies how each array axis maps to mesh axes
# P('data', None) means: shard axis 0 across 'data', replicate axis 1
data_sharding = NamedSharding(mesh, P('data', None))
replicated_sharding = NamedSharding(mesh, P(None, None))
print(f"Mesh: {mesh}")
print(f"Data sharding: {data_sharding}")
print(f"Replicated sharding: {replicated_sharding}")
Mesh: Mesh('data': 1, axis_types=(Auto,))
Data sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec('data', None), memory_kind=device)
Replicated sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec(None, None), memory_kind=device)
# Using sharding for distributed computation
@jit
def matmul_example(x, w):
return x @ w
# Create sharded input and replicated weights
x = jnp.ones((n_devices * 8, 16))
w = jnp.ones((16, 8))
# Place arrays with specific sharding
x_sharded = jax.device_put(x, data_sharding)
w_replicated = jax.device_put(w, replicated_sharding)
print(f"x sharding: {x_sharded.sharding}")
print(f"w sharding: {w_replicated.sharding}")
# JAX automatically handles the distributed computation
result = matmul_example(x_sharded, w_replicated)
print(f"Result shape: {result.shape}")
print(f"Result sharding: {result.sharding}")
x sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec('data', None), memory_kind=device)
w sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec(None, None), memory_kind=device)
Result shape: (8, 8)
Result sharding: NamedSharding(mesh=Mesh('data': 1, axis_types=(Auto,)), spec=PartitionSpec('data', None), memory_kind=device)
# Specifying output sharding with jit
# Use out_shardings to control output placement
@partial(jit, out_shardings=data_sharding)
def compute_with_output_sharding(x):
return jnp.sin(x) + jnp.cos(x)
x = jax.device_put(jnp.ones((n_devices * 4, 8)), data_sharding)
result = compute_with_output_sharding(x)
print(f"Output sharding matches input: {result.sharding == x.sharding}")
Output sharding matches input: True
6. Gradient Accumulation Strategies#
When batch sizes are limited by memory, gradient accumulation allows effective larger batches.
# Gradient accumulation with pmap
def create_model(n_features, n_outputs):
"""Create a simple neural network."""
key = jax.random.PRNGKey(0)
k1, k2 = jax.random.split(key)
return {
'W1': jax.random.normal(k1, (n_features, 64)) * 0.1,
'b1': jnp.zeros(64),
'W2': jax.random.normal(k2, (64, n_outputs)) * 0.1,
'b2': jnp.zeros(n_outputs)
}
def forward(params, x):
"""Forward pass."""
h = jnp.tanh(x @ params['W1'] + params['b1'])
return h @ params['W2'] + params['b2']
def loss_fn(params, x, y):
"""MSE loss."""
pred = forward(params, x)
return jnp.mean((pred - y) ** 2)
# Gradient accumulation function
def accumulate_gradients(params, data_batches):
"""
Accumulate gradients over multiple mini-batches.
data_batches: list of (x, y) tuples
"""
def add_grads(g1, g2):
return jax.tree.map(lambda a, b: a + b, g1, g2)
n_batches = len(data_batches)
# Compute gradients for first batch
x, y = data_batches[0]
total_grads = grad(loss_fn)(params, x, y)
# Accumulate remaining batches
for x, y in data_batches[1:]:
batch_grads = grad(loss_fn)(params, x, y)
total_grads = add_grads(total_grads, batch_grads)
# Average gradients
avg_grads = jax.tree.map(lambda g: g / n_batches, total_grads)
return avg_grads
# Example usage
params = create_model(10, 5)
key = jax.random.PRNGKey(42)
# Create multiple small batches
n_accum_steps = 4
mini_batch_size = 8
data_batches = []
for i in range(n_accum_steps):
key, k1, k2 = jax.random.split(key, 3)
x = jax.random.normal(k1, (mini_batch_size, 10))
y = jax.random.normal(k2, (mini_batch_size, 5))
data_batches.append((x, y))
accumulated_grads = accumulate_gradients(params, data_batches)
print(f"Accumulated over {n_accum_steps} batches of size {mini_batch_size}")
print(f"Effective batch size: {n_accum_steps * mini_batch_size}")
print(f"Gradient W1 norm: {jnp.linalg.norm(accumulated_grads['W1']):.4f}")
Accumulated over 4 batches of size 8
Effective batch size: 32
Gradient W1 norm: 0.4354
# Distributed gradient accumulation with pmap
@partial(pmap, axis_name='devices')
def distributed_accumulate_step(params, x_batches, y_batches):
"""
Each device accumulates gradients over its local batches,
then all-reduce across devices.
x_batches, y_batches: (n_accum_steps, batch_size, features)
"""
n_accum = x_batches.shape[0]
def scan_fn(carry, inputs):
total_grads = carry
x, y = inputs
batch_grads = grad(loss_fn)(params, x, y)
new_total = jax.tree.map(lambda a, b: a + b, total_grads, batch_grads)
return new_total, None
# Initialize with zeros
zero_grads = jax.tree.map(jnp.zeros_like, params)
# Accumulate locally
local_grads, _ = jax.lax.scan(scan_fn, zero_grads, (x_batches, y_batches))
local_grads = jax.tree.map(lambda g: g / n_accum, local_grads)
# All-reduce across devices
global_grads = jax.lax.pmean(local_grads, axis_name='devices')
return global_grads
# Create data for distributed accumulation
# Shape: (n_devices, n_accum_steps, batch_size, features)
n_accum_steps = 4
batch_size = 8
n_features = 10
n_outputs = 5
key = jax.random.PRNGKey(0)
x_all = jax.random.normal(key, (n_devices, n_accum_steps, batch_size, n_features))
y_all = jax.random.normal(jax.random.PRNGKey(1), (n_devices, n_accum_steps, batch_size, n_outputs))
# Replicate params
params = create_model(n_features, n_outputs)
replicated_params = jax.tree.map(lambda x: jnp.stack([x] * n_devices), params)
# Run distributed accumulation
grads = distributed_accumulate_step(replicated_params, x_all, y_all)
effective_batch = n_devices * n_accum_steps * batch_size
print(f"Effective total batch size: {effective_batch}")
print(f"Gradients synchronized across {n_devices} devices")
Effective total batch size: 32
Gradients synchronized across 1 devices
7. Chemical Engineering Example: Parallel Experiment Optimization#
In chemical engineering, we often need to optimize experimental conditions or fit models to data from multiple experiments. Parallel differentiation enables efficient optimization across experiments.
# Parallel experiment fitting for reaction kinetics
def reaction_model(params, conditions):
"""
Arrhenius reaction rate model.
params: [ln_k0, E_over_R, n] (pre-exponential, activation energy, order)
conditions: [T, C] (temperature, concentration)
"""
ln_k0, E_over_R, n = params
T, C = conditions[:, 0], conditions[:, 1]
k = jnp.exp(ln_k0 - E_over_R / T)
rate = k * C ** n
return rate
def experiment_loss(params, conditions, measured_rates):
"""MSE loss for one experiment."""
predicted = reaction_model(params, conditions)
return jnp.mean((predicted - measured_rates) ** 2)
# Generate synthetic experimental data
def generate_experiment(key, true_params, n_points=20):
"""Generate synthetic experimental data with noise."""
k1, k2, k3 = jax.random.split(key, 3)
# Random conditions
T = jax.random.uniform(k1, (n_points,), minval=300, maxval=400)
C = jax.random.uniform(k2, (n_points,), minval=0.1, maxval=2.0)
conditions = jnp.stack([T, C], axis=1)
# True rates with noise
true_rates = reaction_model(true_params, conditions)
noise = jax.random.normal(k3, true_rates.shape) * 0.1 * true_rates
measured_rates = true_rates + noise
return conditions, measured_rates
# True parameters and synthetic experiments
true_params = jnp.array([10.0, 5000.0, 1.5]) # ln_k0, E/R, n
n_experiments = n_devices
experiments = []
for i in range(n_experiments):
key = jax.random.PRNGKey(i * 100)
cond, rates = generate_experiment(key, true_params)
experiments.append((cond, rates))
print(f"Generated {n_experiments} experiments")
print(f"True parameters: ln_k0={true_params[0]:.2f}, E/R={true_params[1]:.0f}, n={true_params[2]:.2f}")
Generated 1 experiments
True parameters: ln_k0=10.00, E/R=5000, n=1.50
# Parallel gradient computation across experiments
@partial(pmap, axis_name='experiments')
def parallel_experiment_gradients(params, conditions, measured_rates):
"""
Compute gradients for each experiment in parallel,
then average across experiments.
"""
local_grads = grad(experiment_loss)(params, conditions, measured_rates)
local_loss = experiment_loss(params, conditions, measured_rates)
# Average gradients and loss across experiments
avg_grads = jax.lax.pmean(local_grads, axis_name='experiments')
avg_loss = jax.lax.pmean(local_loss, axis_name='experiments')
return avg_grads, avg_loss
# Stack experiments for pmap
conditions_stacked = jnp.stack([exp[0] for exp in experiments])
rates_stacked = jnp.stack([exp[1] for exp in experiments])
# Initial parameter guess (replicated)
init_params = jnp.array([8.0, 4000.0, 1.0])
replicated_params = jnp.stack([init_params] * n_experiments)
# Compute gradients
grads, loss = parallel_experiment_gradients(
replicated_params, conditions_stacked, rates_stacked
)
print(f"Initial loss: {loss[0]:.4f}")
print(f"Gradients: {grads[0]}")
Initial loss: 0.0003
Gradients: [ 1.1261558e-03 -2.9959481e-06 -2.7747132e-04]
# Full parallel optimization loop
@partial(pmap, axis_name='experiments')
def parallel_update(params, conditions, measured_rates, learning_rate):
"""Parallel gradient descent step."""
local_grads = grad(experiment_loss)(params, conditions, measured_rates)
local_loss = experiment_loss(params, conditions, measured_rates)
avg_grads = jax.lax.pmean(local_grads, axis_name='experiments')
avg_loss = jax.lax.pmean(local_loss, axis_name='experiments')
new_params = params - learning_rate * avg_grads
return new_params, avg_loss
# Optimization
params = replicated_params
learning_rate = jnp.array([0.1] * n_experiments)
print("Parallel Optimization Across Experiments:")
print("=" * 60)
print(f"{'Iter':<6} {'Loss':<12} {'ln_k0':<10} {'E/R':<10} {'n':<10}")
print("-" * 60)
for i in range(100):
params, loss = parallel_update(
params, conditions_stacked, rates_stacked, learning_rate
)
if i % 20 == 0 or i == 99:
p = params[0] # All devices have same params
print(f"{i:<6} {loss[0]:<12.6f} {p[0]:<10.3f} {p[1]:<10.1f} {p[2]:<10.3f}")
print("-" * 60)
print(f"{'True':<6} {'---':<12} {true_params[0]:<10.3f} {true_params[1]:<10.1f} {true_params[2]:<10.3f}")
Parallel Optimization Across Experiments:
============================================================
Iter Loss ln_k0 E/R n
------------------------------------------------------------
0 0.000272 8.000 4000.0 1.000
20 0.000270 7.998 4000.0 1.001
40 0.000267 7.995 4000.0 1.001
60 0.000264 7.993 4000.0 1.002
80 0.000262 7.991 4000.0 1.002
99 0.000259 7.989 4000.0 1.003
------------------------------------------------------------
True --- 10.000 5000.0 1.500
# Advanced: Parallel flowsheet simulation with different operating conditions
def cstr_steady_state(params, operating_conditions):
"""
CSTR steady-state with kinetic parameters.
params: kinetic parameters [k0, E_over_R]
operating_conditions: [T, tau, C_in] for each experiment
"""
k0, E_over_R = params
T, tau, C_in = operating_conditions
k = k0 * jnp.exp(-E_over_R / T)
# Steady-state: C = C_in / (1 + k*tau)
C_out = C_in / (1 + k * tau)
conversion = 1 - C_out / C_in
return conversion
@partial(pmap, axis_name='conditions')
def parallel_cstr_sensitivity(params, operating_conditions):
"""
Compute sensitivity of conversion to kinetic parameters
for multiple operating conditions in parallel.
"""
# Conversion and its gradient
conversion = cstr_steady_state(params, operating_conditions)
sensitivity = grad(cstr_steady_state)(params, operating_conditions)
# Normalized sensitivity
norm_sens = (params / conversion) * sensitivity
return conversion, norm_sens
# Different operating conditions to analyze
kinetic_params = jnp.array([1e6, 5000.0]) # k0, E/R
# Generate range of conditions (T, tau, C_in)
T_values = jnp.linspace(320, 380, n_devices)
tau_values = jnp.full(n_devices, 10.0)
C_in_values = jnp.full(n_devices, 1.0)
operating_conditions = jnp.stack([T_values, tau_values, C_in_values], axis=1)
replicated_params = jnp.stack([kinetic_params] * n_devices)
# Parallel computation
conversions, sensitivities = parallel_cstr_sensitivity(
replicated_params, operating_conditions
)
print("Parallel CSTR Sensitivity Analysis:")
print("=" * 60)
print(f"{'T (K)':<10} {'Conv.':<10} {'Sens(k0)':<15} {'Sens(E/R)':<15}")
print("-" * 60)
for i in range(n_devices):
print(f"{T_values[i]:<10.1f} {conversions[i]:<10.4f} "
f"{sensitivities[i, 0]:<15.4f} {sensitivities[i, 1]:<15.4f}")
Parallel CSTR Sensitivity Analysis:
============================================================
T (K) Conv. Sens(k0) Sens(E/R)
------------------------------------------------------------
320.0 0.6208 0.3792 -5.9244
Summary#
This tutorial covered distributed and parallel differentiation in JAX:
Tool |
Use Case |
Key Features |
|---|---|---|
|
Data parallelism |
Simple API, collective ops (pmean, psum) |
|
Flexible array distribution |
NamedSharding, SingleDeviceSharding |
|
Named device topology |
Multi-dimensional device grids |
|
Explicit SPMD |
Fine-grained control over sharding |
Key Patterns for Distributed Gradients#
Data parallel training:
@partial(pmap, axis_name='devices') def train_step(params, x, y): grads = grad(loss)(params, x, y) return jax.lax.pmean(grads, 'devices')
Gradient accumulation: Accumulate locally, then all-reduce
Parallel experiments: Same model, different data shards
Chemical Engineering Applications#
Parameter estimation: Fit kinetic models to multiple experiments in parallel
Sensitivity analysis: Analyze many operating conditions simultaneously
Optimization: Large-scale process optimization with distributed gradients
Uncertainty propagation: Parallel Monte Carlo sampling