Tutorial 26: GPU Acceleration for Differentiable Programming#
This tutorial demonstrates how to leverage GPU acceleration with JAX for differentiable programming. We’ll cover:
Detecting and selecting devices (CPU vs GPU)
Moving computations between devices
Performance comparison: CPU vs GPU
Best practices for GPU-accelerated autodiff
Memory management on GPU
Chemical engineering example: Large-scale optimization
import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, jacfwd, jacrev
import numpy as np
import time
# Enable 64-bit precision
jax.config.update("jax_enable_x64", True)
print(f"JAX version: {jax.__version__}")
JAX version: 0.8.2
1. Device Detection and Selection#
JAX automatically detects available accelerators. Let’s explore what devices are available.
# List all available devices
all_devices = jax.devices()
print(f"All devices: {all_devices}")
print(f"Number of devices: {len(all_devices)}")
# Check for specific device types
cpu_devices = jax.devices('cpu')
print(f"\nCPU devices: {cpu_devices}")
try:
gpu_devices = jax.devices('gpu')
print(f"GPU devices: {gpu_devices}")
HAS_GPU = True
except RuntimeError:
print("No GPU devices available")
HAS_GPU = False
# Default device (where arrays are created by default)
print(f"\nDefault backend: {jax.default_backend()}")
WARNING:2026-01-10 21:03:17,900: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.
All devices: [CpuDevice(id=0)]
Number of devices: 1
CPU devices: [CpuDevice(id=0)]
No GPU devices available
Default backend: cpu
# Get detailed device information
for i, device in enumerate(jax.devices()):
print(f"Device {i}:")
print(f" Platform: {device.platform}")
print(f" Device kind: {device.device_kind}")
print(f" ID: {device.id}")
Device 0:
Platform: cpu
Device kind: cpu
ID: 0
2. Placing Arrays on Specific Devices#
Use jax.device_put() to explicitly place arrays on specific devices.
# Create an array (goes to default device)
x_default = jnp.ones((1000, 1000))
print(f"Default placement: {x_default.devices()}")
# Explicitly place on CPU
cpu = jax.devices('cpu')[0]
x_cpu = jax.device_put(x_default, cpu)
print(f"CPU placement: {x_cpu.devices()}")
# Explicitly place on GPU (if available)
if HAS_GPU:
gpu = jax.devices('gpu')[0]
x_gpu = jax.device_put(x_default, gpu)
print(f"GPU placement: {x_gpu.devices()}")
Default placement: {CpuDevice(id=0)}
CPU placement: {CpuDevice(id=0)}
# Arrays stay on their device through operations
if HAS_GPU:
gpu = jax.devices('gpu')[0]
a_gpu = jax.device_put(jnp.array([1.0, 2.0, 3.0]), gpu)
b_gpu = jax.device_put(jnp.array([4.0, 5.0, 6.0]), gpu)
# Operations on GPU arrays produce GPU arrays
c_gpu = a_gpu + b_gpu
print(f"Result device: {c_gpu.devices()}")
print(f"Result: {c_gpu}")
3. Performance Comparison: CPU vs GPU#
Let’s compare performance for different operations and problem sizes.
def benchmark(fn, *args, n_warmup=3, n_runs=10):
"""Benchmark a function with warmup runs."""
# Warmup
for _ in range(n_warmup):
result = fn(*args)
if hasattr(result, 'block_until_ready'):
result.block_until_ready()
# Timed runs
times = []
for _ in range(n_runs):
start = time.perf_counter()
result = fn(*args)
if hasattr(result, 'block_until_ready'):
result.block_until_ready()
times.append(time.perf_counter() - start)
return np.mean(times), np.std(times), result
3.1 Matrix Multiplication#
if HAS_GPU:
print("Matrix Multiplication: CPU vs GPU")
print("=" * 60)
cpu = jax.devices('cpu')[0]
gpu = jax.devices('gpu')[0]
sizes = [256, 512, 1024, 2048]
for n in sizes:
# Create matrices on each device
key = jax.random.PRNGKey(0)
A = jax.random.normal(key, (n, n))
B = jax.random.normal(jax.random.PRNGKey(1), (n, n))
A_cpu = jax.device_put(A, cpu)
B_cpu = jax.device_put(B, cpu)
A_gpu = jax.device_put(A, gpu)
B_gpu = jax.device_put(B, gpu)
# JIT-compiled matmul
matmul = jit(lambda x, y: x @ y)
# Benchmark
cpu_time, cpu_std, _ = benchmark(matmul, A_cpu, B_cpu)
gpu_time, gpu_std, _ = benchmark(matmul, A_gpu, B_gpu)
speedup = cpu_time / gpu_time
print(f"Size {n}x{n}: CPU={cpu_time*1000:.2f}ms, GPU={gpu_time*1000:.2f}ms, Speedup={speedup:.1f}x")
else:
print("GPU not available - skipping comparison")
GPU not available - skipping comparison
3.2 Gradient Computation#
if HAS_GPU:
print("Gradient Computation: CPU vs GPU")
print("=" * 60)
def loss_fn(W, x, y):
"""Neural network loss: MSE of single layer."""
pred = jnp.tanh(x @ W)
return jnp.mean((pred - y) ** 2)
grad_fn = jit(grad(loss_fn))
sizes = [(500, 256), (1000, 512), (2000, 512), (4000, 512)]
for n_samples, n_features in sizes:
key = jax.random.PRNGKey(0)
k1, k2, k3 = jax.random.split(key, 3)
W = jax.random.normal(k1, (n_features, n_features))
x = jax.random.normal(k2, (n_samples, n_features))
y = jax.random.normal(k3, (n_samples, n_features))
# CPU
W_cpu = jax.device_put(W, cpu)
x_cpu = jax.device_put(x, cpu)
y_cpu = jax.device_put(y, cpu)
# GPU
W_gpu = jax.device_put(W, gpu)
x_gpu = jax.device_put(x, gpu)
y_gpu = jax.device_put(y, gpu)
cpu_time, _, _ = benchmark(grad_fn, W_cpu, x_cpu, y_cpu)
gpu_time, _, _ = benchmark(grad_fn, W_gpu, x_gpu, y_gpu)
speedup = cpu_time / gpu_time
print(f"W: {n_features}x{n_features}, batch: {n_samples} -> "
f"CPU={cpu_time*1000:.2f}ms, GPU={gpu_time*1000:.2f}ms, Speedup={speedup:.1f}x")
else:
print("GPU not available - skipping comparison")
GPU not available - skipping comparison
3.3 Jacobian Computation#
if HAS_GPU:
print("Jacobian Computation: CPU vs GPU")
print("=" * 60)
def vector_fn(x):
"""A vector-valued function."""
return jnp.tanh(x @ x.T).sum(axis=1)
jacobian_fn = jit(jacrev(vector_fn))
sizes = [32, 64, 128, 256]
for n in sizes:
key = jax.random.PRNGKey(0)
x = jax.random.normal(key, (n, n))
x_cpu = jax.device_put(x, cpu)
x_gpu = jax.device_put(x, gpu)
cpu_time, _, J_cpu = benchmark(jacobian_fn, x_cpu, n_runs=5)
gpu_time, _, J_gpu = benchmark(jacobian_fn, x_gpu, n_runs=5)
speedup = cpu_time / gpu_time
print(f"Input: {n}x{n}, Jacobian: {J_cpu.shape} -> "
f"CPU={cpu_time*1000:.1f}ms, GPU={gpu_time*1000:.1f}ms, Speedup={speedup:.1f}x")
else:
print("GPU not available - skipping comparison")
GPU not available - skipping comparison
3.4 Batched Operations with vmap#
if HAS_GPU:
print("Batched Gradient (vmap): CPU vs GPU")
print("=" * 60)
def single_loss(w, x, y):
"""Loss for a single sample."""
pred = jnp.tanh(jnp.dot(x, w))
return jnp.sum((pred - y) ** 2)
# Batch over samples
batched_grad = jit(vmap(grad(single_loss), in_axes=(None, 0, 0)))
n_features = 128
batch_sizes = [100, 500, 1000, 2000, 5000]
for batch_size in batch_sizes:
key = jax.random.PRNGKey(0)
k1, k2, k3 = jax.random.split(key, 3)
w = jax.random.normal(k1, (n_features, n_features))
x = jax.random.normal(k2, (batch_size, n_features))
y = jax.random.normal(k3, (batch_size, n_features))
w_cpu, x_cpu, y_cpu = [jax.device_put(a, cpu) for a in (w, x, y)]
w_gpu, x_gpu, y_gpu = [jax.device_put(a, gpu) for a in (w, x, y)]
cpu_time, _, _ = benchmark(batched_grad, w_cpu, x_cpu, y_cpu)
gpu_time, _, _ = benchmark(batched_grad, w_gpu, x_gpu, y_gpu)
speedup = cpu_time / gpu_time
print(f"Batch size {batch_size:>5}: CPU={cpu_time*1000:.2f}ms, "
f"GPU={gpu_time*1000:.2f}ms, Speedup={speedup:.1f}x")
else:
print("GPU not available - skipping comparison")
GPU not available - skipping comparison
4. Best Practices for GPU-Accelerated Autodiff#
4.1 Always Use JIT Compilation#
JIT is essential for GPU performance - without it, each operation incurs Python overhead.
if HAS_GPU:
print("JIT vs No-JIT on GPU")
print("=" * 60)
def complex_fn(x):
"""Multiple operations - benefits greatly from JIT."""
for _ in range(10):
x = jnp.sin(x) + jnp.cos(x)
x = x @ x.T
x = jnp.tanh(x)
return x.sum()
grad_no_jit = grad(complex_fn)
grad_with_jit = jit(grad(complex_fn))
x = jax.random.normal(jax.random.PRNGKey(0), (100, 100))
x_gpu = jax.device_put(x, gpu)
# Benchmark (fewer runs for non-JIT since it's slow)
no_jit_time, _, _ = benchmark(grad_no_jit, x_gpu, n_warmup=1, n_runs=3)
jit_time, _, _ = benchmark(grad_with_jit, x_gpu, n_warmup=3, n_runs=10)
print(f"Without JIT: {no_jit_time*1000:.2f}ms")
print(f"With JIT: {jit_time*1000:.2f}ms")
print(f"Speedup: {no_jit_time/jit_time:.1f}x")
else:
print("GPU not available - skipping")
GPU not available - skipping
4.2 Minimize Host-Device Transfers#
Transferring data between CPU and GPU is expensive. Keep data on the GPU throughout computation.
if HAS_GPU:
print("Data Transfer Overhead")
print("=" * 60)
@jit
def compute(x):
return jnp.sin(x) + jnp.cos(x)
sizes = [1000, 10000, 100000, 1000000]
for n in sizes:
x_np = np.random.randn(n)
x_gpu = jax.device_put(jnp.array(x_np), gpu)
# Time: transfer + compute + transfer back
def with_transfer():
x_gpu_local = jax.device_put(jnp.array(x_np), gpu)
result = compute(x_gpu_local)
return np.array(result) # Transfer back to CPU
# Time: compute only (data already on GPU)
def no_transfer():
return compute(x_gpu)
transfer_time, _, _ = benchmark(with_transfer, n_runs=20)
no_transfer_time, _, _ = benchmark(no_transfer, n_runs=20)
overhead = (transfer_time - no_transfer_time) / no_transfer_time * 100
print(f"Size {n:>7}: With transfer={transfer_time*1000:.3f}ms, "
f"No transfer={no_transfer_time*1000:.3f}ms, Overhead={overhead:.0f}%")
else:
print("GPU not available - skipping")
GPU not available - skipping
4.3 Use Appropriate Batch Sizes#
GPUs excel at parallel computation. Too-small batches underutilize the GPU.
if HAS_GPU:
print("Optimal Batch Sizes")
print("=" * 60)
def batch_matmul(A, B):
"""Batched matrix multiplication."""
return jnp.einsum('bij,bjk->bik', A, B)
batch_fn = jit(batch_matmul)
matrix_size = 64
batch_sizes = [1, 10, 50, 100, 500, 1000, 2000]
print(f"Matrix size: {matrix_size}x{matrix_size}")
print("-" * 60)
for batch_size in batch_sizes:
key = jax.random.PRNGKey(0)
A = jax.random.normal(key, (batch_size, matrix_size, matrix_size))
B = jax.random.normal(jax.random.PRNGKey(1), (batch_size, matrix_size, matrix_size))
A_gpu = jax.device_put(A, gpu)
B_gpu = jax.device_put(B, gpu)
time_taken, _, _ = benchmark(batch_fn, A_gpu, B_gpu)
throughput = batch_size / time_taken
print(f"Batch {batch_size:>5}: {time_taken*1000:.3f}ms, "
f"Throughput={throughput:.0f} matmuls/sec")
else:
print("GPU not available - skipping")
GPU not available - skipping
5. Memory Management on GPU#
GPU memory is limited. JAX provides tools to manage it effectively.
if HAS_GPU:
# Check current memory usage
# Note: This requires jax >= 0.4.1
try:
memory_stats = jax.devices('gpu')[0].memory_stats()
print("GPU Memory Stats:")
for key, value in memory_stats.items():
if 'bytes' in key:
print(f" {key}: {value / 1e9:.2f} GB")
except Exception as e:
print(f"Memory stats not available: {e}")
# Gradient checkpointing to reduce memory usage
from jax import checkpoint
def memory_heavy_fn(x):
"""Function with many intermediate values."""
for _ in range(10):
x = jnp.tanh(x @ x.T + x)
return x.sum()
# Without checkpointing: stores all intermediates
grad_no_checkpoint = jit(grad(memory_heavy_fn))
# With checkpointing: recomputes intermediates during backward pass
grad_with_checkpoint = jit(grad(checkpoint(memory_heavy_fn)))
if HAS_GPU:
x = jax.random.normal(jax.random.PRNGKey(0), (200, 200))
x_gpu = jax.device_put(x, gpu)
# Both produce the same result
result1 = grad_no_checkpoint(x_gpu)
result2 = grad_with_checkpoint(x_gpu)
print(f"Results match: {jnp.allclose(result1, result2)}")
print("\nCheckpointing trades compute for memory:")
print(" - Without: Faster, but uses more GPU memory")
print(" - With: Slower, but uses less GPU memory")
6. Chemical Engineering Example: Large-Scale Optimization#
Let’s apply GPU acceleration to a realistic chemical engineering problem: optimizing reaction conditions across a large parameter space.
# Multi-reaction kinetic model with Arrhenius kinetics
def kinetic_model_loss(params, conditions, measurements):
"""
Parameter estimation loss with multi-reaction kinetics.
params: (n_reactions, 2) - [ln_k0, E/R] for each reaction
conditions: (n_experiments, 2) - [T, tau] for each experiment
measurements: (n_experiments, n_species) - measured concentrations
Returns: scalar loss
"""
T = conditions[:, 0:1] # (n_exp, 1)
tau = conditions[:, 1:2] # (n_exp, 1)
# Compute rate constants for all reactions
ln_k0 = params[:, 0] # (n_reactions,)
E_over_R = params[:, 1] # (n_reactions,)
# k = k0 * exp(-E/RT) for each reaction and experiment
k = jnp.exp(ln_k0[None, :] - E_over_R[None, :] / T) # (n_exp, n_reactions)
# Simple reaction network: A -> B -> C
# Species concentrations at steady state (simplified)
C_A = 1.0 / (1 + k[:, 0:1] * tau)
C_B = k[:, 0:1] * C_A * tau / (1 + k[:, 1:2] * tau)
C_C = 1 - C_A - C_B
predictions = jnp.concatenate([C_A, C_B, C_C], axis=1)
return jnp.mean((predictions - measurements) ** 2)
def run_optimization(device, n_experiments, n_iterations=100):
"""Run parameter estimation optimization on specified device."""
key = jax.random.PRNGKey(42)
k1, k2, k3 = jax.random.split(key, 3)
# True parameters
true_params = jnp.array([
[10.0, 5000.0], # Reaction 1: ln_k0, E/R
[12.0, 6000.0], # Reaction 2
])
# Generate synthetic experiment data
T = jax.random.uniform(k1, (n_experiments, 1), minval=300, maxval=400)
tau = jax.random.uniform(k2, (n_experiments, 1), minval=1, maxval=10)
conditions = jnp.concatenate([T, tau], axis=1)
# Generate true concentrations with noise
measurements = jax.random.normal(k3, (n_experiments, 3)) * 0.01 + 0.33
# Initial guess
params = jnp.array([
[8.0, 4000.0],
[10.0, 5000.0],
])
# Move to device
params = jax.device_put(params, device)
conditions = jax.device_put(conditions, device)
measurements = jax.device_put(measurements, device)
# JIT compile gradient function
loss_and_grad = jit(jax.value_and_grad(kinetic_model_loss))
# Optimization loop
learning_rate = 0.1
start_time = time.perf_counter()
for i in range(n_iterations):
loss, grads = loss_and_grad(params, conditions, measurements)
params = params - learning_rate * grads
# Ensure computation is complete (for accurate timing)
if i == n_iterations - 1:
loss.block_until_ready()
total_time = time.perf_counter() - start_time
return total_time, float(loss)
if HAS_GPU:
print("Large-Scale Parameter Estimation: CPU vs GPU")
print("=" * 70)
cpu = jax.devices('cpu')[0]
gpu = jax.devices('gpu')[0]
experiment_sizes = [1000, 5000, 10000, 50000, 100000]
n_iterations = 100
print(f"Running {n_iterations} optimization iterations for each size")
print("-" * 70)
print(f"{'Experiments':<12} {'CPU Time':<12} {'GPU Time':<12} {'Speedup':<10} {'Final Loss'}")
print("-" * 70)
for n_exp in experiment_sizes:
cpu_time, cpu_loss = run_optimization(cpu, n_exp, n_iterations)
gpu_time, gpu_loss = run_optimization(gpu, n_exp, n_iterations)
speedup = cpu_time / gpu_time
print(f"{n_exp:<12} {cpu_time:<12.3f} {gpu_time:<12.3f} {speedup:<10.1f}x {gpu_loss:.6f}")
else:
print("GPU not available - running on CPU only")
cpu = jax.devices('cpu')[0]
cpu_time, cpu_loss = run_optimization(cpu, 10000, 100)
print(f"CPU time: {cpu_time:.3f}s, Final loss: {cpu_loss:.6f}")
GPU not available - running on CPU only
CPU time: 0.366s, Final loss: 0.112481
6.2 Parallel Sensitivity Analysis#
# Compute sensitivities across many operating conditions in parallel
def cstr_conversion(kinetic_params, operating_conditions):
"""
CSTR conversion model.
kinetic_params: [k0, E_R, n] - kinetic parameters
operating_conditions: [T, C_in, tau] - operating conditions
"""
k0, E_R, n = kinetic_params
T, C_in, tau = operating_conditions
k = k0 * jnp.exp(-E_R / T)
# CSTR conversion for nth order reaction
C_out = C_in / (1 + k * tau * C_in ** (n - 1))
conversion = (C_in - C_out) / C_in
return conversion
# Compute Jacobian of conversion w.r.t. both params and conditions
@jit
def compute_sensitivities(kinetic_params, operating_conditions):
"""Compute all sensitivities for CSTR conversion."""
# Sensitivity to kinetic parameters
d_conv_d_params = jacrev(cstr_conversion, argnums=0)(
kinetic_params, operating_conditions
)
# Sensitivity to operating conditions
d_conv_d_conditions = jacrev(cstr_conversion, argnums=1)(
kinetic_params, operating_conditions
)
return d_conv_d_params, d_conv_d_conditions
# Batch over many conditions
batched_sensitivity = jit(vmap(
compute_sensitivities,
in_axes=(None, 0) # Same params, different conditions
))
if HAS_GPU:
print("Parallel Sensitivity Analysis: CPU vs GPU")
print("=" * 60)
kinetic_params = jnp.array([1e6, 5000.0, 1.5]) # k0, E/R, n
n_conditions_list = [100, 1000, 10000, 50000, 100000]
for n_conditions in n_conditions_list:
key = jax.random.PRNGKey(0)
k1, k2, k3 = jax.random.split(key, 3)
# Random operating conditions
T = jax.random.uniform(k1, (n_conditions,), minval=300, maxval=400)
C_in = jax.random.uniform(k2, (n_conditions,), minval=0.5, maxval=2.0)
tau = jax.random.uniform(k3, (n_conditions,), minval=1, maxval=20)
operating_conditions = jnp.stack([T, C_in, tau], axis=1)
# CPU
params_cpu = jax.device_put(kinetic_params, cpu)
cond_cpu = jax.device_put(operating_conditions, cpu)
# GPU
params_gpu = jax.device_put(kinetic_params, gpu)
cond_gpu = jax.device_put(operating_conditions, gpu)
cpu_time, _, _ = benchmark(batched_sensitivity, params_cpu, cond_cpu)
gpu_time, _, _ = benchmark(batched_sensitivity, params_gpu, cond_gpu)
speedup = cpu_time / gpu_time
print(f"{n_conditions:>6} conditions: CPU={cpu_time*1000:.2f}ms, "
f"GPU={gpu_time*1000:.2f}ms, Speedup={speedup:.1f}x")
else:
print("GPU not available - skipping comparison")
GPU not available - skipping comparison
7. When to Use GPU vs CPU#
GPU Advantages#
Large batch sizes: Parallel processing of many samples
Large matrices: Matrix operations (matmul, einsum)
Deep computation graphs: Many layers/operations to differentiate through
Parameter sweeps: Same computation over many parameter values
CPU May Be Better For#
Small problems: Transfer overhead exceeds compute savings
Sequential dependencies: Limited parallelism
Memory-bound: Large intermediate results that don’t fit on GPU
Sparse operations: Irregular memory access patterns
# Example: Small problem where CPU might win
if HAS_GPU:
print("Small Problem Performance")
print("=" * 60)
def small_fn(x):
return jnp.sum(x ** 2)
grad_fn = jit(grad(small_fn))
sizes = [10, 50, 100, 500, 1000]
for n in sizes:
x = jax.random.normal(jax.random.PRNGKey(0), (n,))
x_cpu = jax.device_put(x, cpu)
x_gpu = jax.device_put(x, gpu)
cpu_time, _, _ = benchmark(grad_fn, x_cpu, n_runs=100)
gpu_time, _, _ = benchmark(grad_fn, x_gpu, n_runs=100)
winner = "CPU" if cpu_time < gpu_time else "GPU"
speedup = max(cpu_time, gpu_time) / min(cpu_time, gpu_time)
print(f"Size {n:>4}: CPU={cpu_time*1000:.3f}ms, GPU={gpu_time*1000:.3f}ms -> {winner} wins ({speedup:.1f}x)")
Summary#
Topic |
Key Points |
|---|---|
Device Selection |
Use |
JIT Compilation |
Essential for GPU performance - always JIT your hot paths |
Data Transfer |
Minimize CPU-GPU transfers; keep data on GPU through computation |
Batch Size |
Larger batches better utilize GPU parallelism |
Memory |
Use |
When GPU Wins |
Large matrices, many parallel operations, big batches |
When CPU Wins |
Small problems, sequential dependencies, sparse operations |
Chemical Engineering Applications#
GPUs excel at:
Parameter estimation with many experiments
Sensitivity analysis across many conditions
Monte Carlo uncertainty quantification
Optimization with large parameter spaces
Surrogate models (neural networks) for process simulation