JAX Fundamentals: Differentiable Programming#
This notebook introduces the core concepts of differentiable programming with JAX.
What is Differentiable Programming?#
Differentiable programming treats programs as mathematical functions that can be differentiated automatically. This enables:
Optimization: Find parameters that minimize/maximize objectives
Sensitivity analysis: Understand how outputs depend on inputs
Machine learning: Train models via gradient descent
Scientific computing: Solve inverse problems, ODEs, PDEs
Why JAX?#
JAX provides:
Automatic differentiation (
grad,jacfwd,jacrev,hessian)JIT compilation for speed (
jit)Vectorization (
vmap)Functional programming paradigm
NumPy-compatible API
import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, jacfwd, jacrev, hessian
import numpy as np
# Enable 64-bit precision (important for scientific computing)
jax.config.update("jax_enable_x64", True)
print(f"JAX version: {jax.__version__}")
print(f"Devices: {jax.devices()}")
WARNING:2026-01-10 20:21:44,752: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.
JAX version: 0.8.2
Devices: [CpuDevice(id=0)]
1. Basic Automatic Differentiation#
1.1 Scalar Functions#
The grad function computes the gradient of a scalar-valued function.
# Define a simple function
def f(x):
return x**3 - 2*x**2 + x
# Compute its derivative: f'(x) = 3x² - 4x + 1
df = grad(f)
# Evaluate at x = 2.0
x = 2.0
print(f"f({x}) = {f(x)}")
print(f"f'({x}) = {df(x)}")
print(f"Expected: 3*{x}² - 4*{x} + 1 = {3*x**2 - 4*x + 1}")
f(2.0) = 2.0
f'(2.0) = 5.0
Expected: 3*2.0² - 4*2.0 + 1 = 5.0
# Higher-order derivatives
d2f = grad(grad(f)) # Second derivative: f''(x) = 6x - 4
d3f = grad(grad(grad(f))) # Third derivative: f'''(x) = 6
print(f"f''({x}) = {d2f(x)} (expected: {6*x - 4})")
print(f"f'''({x}) = {d3f(x)} (expected: 6)")
f''(2.0) = 8.0 (expected: 8.0)
f'''(2.0) = 6.0 (expected: 6)
1.2 Multivariate Functions#
For functions of multiple variables, grad returns the gradient vector.
def g(params):
"""g(x, y) = x² + xy + y²"""
x, y = params[0], params[1]
return x**2 + x*y + y**2
# Gradient: ∇g = [2x + y, x + 2y]
grad_g = grad(g)
point = jnp.array([1.0, 2.0])
gradient = grad_g(point)
print(f"Point: {point}")
print(f"g(x,y) = {g(point)}")
print(f"∇g = {gradient}")
print(f"Expected: [{2*1 + 2}, {1 + 2*2}] = [4, 5]")
Point: [1. 2.]
g(x,y) = 7.0
∇g = [4. 5.]
Expected: [4, 5] = [4, 5]
1.3 Specifying Which Argument to Differentiate#
Use argnums to specify which argument(s) to differentiate.
def h(x, y, z):
"""h(x, y, z) = x*y + y*z + z*x"""
return x*y + y*z + z*x
# Partial derivatives
dh_dx = grad(h, argnums=0) # ∂h/∂x = y + z
dh_dy = grad(h, argnums=1) # ∂h/∂y = x + z
dh_dz = grad(h, argnums=2) # ∂h/∂z = y + x
# All at once
dh_all = grad(h, argnums=(0, 1, 2))
x, y, z = 1.0, 2.0, 3.0
print(f"∂h/∂x = {dh_dx(x, y, z)} (expected: y+z = {y+z})")
print(f"∂h/∂y = {dh_dy(x, y, z)} (expected: x+z = {x+z})")
print(f"∂h/∂z = {dh_dz(x, y, z)} (expected: x+y = {x+y})")
print(f"\nAll gradients: {dh_all(x, y, z)}")
∂h/∂x = 5.0 (expected: y+z = 5.0)
∂h/∂y = 4.0 (expected: x+z = 4.0)
∂h/∂z = 3.0 (expected: x+y = 3.0)
All gradients: (Array(5., dtype=float64, weak_type=True), Array(4., dtype=float64, weak_type=True), Array(3., dtype=float64, weak_type=True))
2. Jacobians and Hessians#
2.1 Jacobian Matrix#
For vector-valued functions \(f: \mathbb{R}^n \to \mathbb{R}^m\), the Jacobian is the \(m \times n\) matrix:
def vector_fn(x):
"""f: R² → R³"""
return jnp.array([
x[0]**2 + x[1], # f₁ = x² + y
x[0] * x[1], # f₂ = xy
jnp.sin(x[0] + x[1]) # f₃ = sin(x+y)
])
# Two ways to compute Jacobian:
# - jacfwd: Forward-mode (efficient when n < m)
# - jacrev: Reverse-mode (efficient when n > m)
x = jnp.array([1.0, 2.0])
J_fwd = jacfwd(vector_fn)(x)
J_rev = jacrev(vector_fn)(x)
print("Jacobian (3×2 matrix):")
print(J_fwd)
print(f"\nJ[0,:] = ∂f₁/∂x = [{2*x[0]}, 1] (∂(x²+y)/∂x, ∂(x²+y)/∂y)")
print(f"J[1,:] = ∂f₂/∂x = [{x[1]}, {x[0]}] (∂(xy)/∂x, ∂(xy)/∂y)")
Jacobian (3×2 matrix):
[[ 2. 1. ]
[ 2. 1. ]
[-0.9899925 -0.9899925]]
J[0,:] = ∂f₁/∂x = [2.0, 1] (∂(x²+y)/∂x, ∂(x²+y)/∂y)
J[1,:] = ∂f₂/∂x = [2.0, 1.0] (∂(xy)/∂x, ∂(xy)/∂y)
2.2 Hessian Matrix#
For scalar-valued functions, the Hessian is the matrix of second derivatives:
def quadratic(x):
"""f(x,y) = x² + 3xy + 2y²"""
return x[0]**2 + 3*x[0]*x[1] + 2*x[1]**2
# Hessian: [[2, 3], [3, 4]]
H = hessian(quadratic)
x = jnp.array([1.0, 1.0])
print("Hessian matrix:")
print(H(x))
print("\nExpected: [[∂²f/∂x², ∂²f/∂x∂y], [∂²f/∂y∂x, ∂²f/∂y²]] = [[2, 3], [3, 4]]")
Hessian matrix:
[[2. 3.]
[3. 4.]]
Expected: [[∂²f/∂x², ∂²f/∂x∂y], [∂²f/∂y∂x, ∂²f/∂y²]] = [[2, 3], [3, 4]]
2.3 Forward vs Reverse Mode: When to Use jacfwd vs jacrev#
Both compute the same Jacobian, but with different computational costs:
Mode |
Function |
Cost |
Best When |
|---|---|---|---|
Forward |
|
O(n) Jacobian-vector products |
n < m (few inputs, many outputs) |
Reverse |
|
O(m) vector-Jacobian products |
n > m (many inputs, few outputs) |
Where \(f: \mathbb{R}^n \to \mathbb{R}^m\).
Rule of thumb:
jacrevfor loss functions (many params → scalar): likegradjacfwdfor expanding functions (few inputs → many outputs)
# Example: Few inputs → Many outputs (jacfwd is better)
def few_to_many(x):
"""f: R² → R¹⁰⁰"""
return jnp.array([x[0]**i + x[1]**(i+1) for i in range(1, 101)])
x_small = jnp.array([1.5, 0.9])
# Time both approaches
import time
# Warm up
_ = jacfwd(few_to_many)(x_small)
_ = jacrev(few_to_many)(x_small)
start = time.time()
for _ in range(100):
J_fwd = jacfwd(few_to_many)(x_small)
time_fwd = time.time() - start
start = time.time()
for _ in range(100):
J_rev = jacrev(few_to_many)(x_small)
time_rev = time.time() - start
print(f"f: R² → R¹⁰⁰ (few inputs, many outputs)")
print(f" jacfwd: {time_fwd*1000:.2f} ms")
print(f" jacrev: {time_rev*1000:.2f} ms")
print(f" → jacfwd is {time_rev/time_fwd:.1f}x faster (as expected)")
print(f"\nJacobian shape: {J_fwd.shape}")
f: R² → R¹⁰⁰ (few inputs, many outputs)
jacfwd: 34641.99 ms
jacrev: 77938.65 ms
→ jacfwd is 2.2x faster (as expected)
Jacobian shape: (100, 2)
# Example: Many inputs → Few outputs (jacrev is better)
def many_to_few(x):
"""f: R¹⁰⁰ → R²"""
return jnp.array([jnp.sum(x**2), jnp.prod(jnp.tanh(x))])
x_large = jnp.ones(100) * 0.5
# Warm up
_ = jacfwd(many_to_few)(x_large)
_ = jacrev(many_to_few)(x_large)
start = time.time()
for _ in range(100):
J_fwd = jacfwd(many_to_few)(x_large)
time_fwd = time.time() - start
start = time.time()
for _ in range(100):
J_rev = jacrev(many_to_few)(x_large)
time_rev = time.time() - start
print(f"f: R¹⁰⁰ → R² (many inputs, few outputs)")
print(f" jacfwd: {time_fwd*1000:.2f} ms")
print(f" jacrev: {time_rev*1000:.2f} ms")
print(f" → jacrev is {time_fwd/time_rev:.1f}x faster (as expected)")
print(f"\nJacobian shape: {J_rev.shape}")
f: R¹⁰⁰ → R² (many inputs, few outputs)
jacfwd: 485.14 ms
jacrev: 2563.15 ms
→ jacrev is 0.2x faster (as expected)
Jacobian shape: (2, 100)
2.4 VJP and JVP: Vector-Jacobian and Jacobian-Vector Products#
Sometimes you don’t need the full Jacobian—just its product with a vector:
JVP (Jacobian-Vector Product): \(J \cdot v\) — “pushforward” tangent vectors
VJP (Vector-Jacobian Product): \(v^T \cdot J\) — “pullback” cotangent vectors
These are O(1) in the Jacobian dimensions (vs O(n×m) for full Jacobian).
from jax import jvp, vjp
def f(x):
"""f: R³ → R²"""
return jnp.array([x[0]*x[1], x[1]*x[2] + x[0]**2])
x = jnp.array([1.0, 2.0, 3.0])
# JVP: Compute f(x) and J·v in one forward pass
# Useful for: directional derivatives, forward-mode AD
v = jnp.array([1.0, 0.0, 0.0]) # tangent vector (direction)
primals, tangents = jvp(f, (x,), (v,))
print("JVP (Jacobian-Vector Product):")
print(f" f(x) = {primals}")
print(f" J·v = {tangents} (directional derivative in direction v)")
# Verify: J·v should equal the first column of J (since v = [1,0,0])
J = jacfwd(f)(x)
print(f" J @ v = {J @ v} (verification)")
print(f"\nFull Jacobian:\n{J}")
JVP (Jacobian-Vector Product):
f(x) = [2. 7.]
J·v = [2. 2.] (directional derivative in direction v)
J @ v = [2. 2.] (verification)
Full Jacobian:
[[2. 1. 0.]
[2. 3. 2.]]
# VJP: Compute f(x) and function to compute vᵀ·J
# Useful for: backpropagation, reverse-mode AD, computing gradients
primals, vjp_fn = vjp(f, x)
print("VJP (Vector-Jacobian Product):")
print(f" f(x) = {primals}")
# The vjp_fn takes a cotangent vector (same shape as output)
cotangent = jnp.array([1.0, 0.0]) # "how much we care about each output"
grad_x, = vjp_fn(cotangent)
print(f" vᵀ·J = {grad_x} (gradient w.r.t. x when caring only about f₁)")
# Verify: vᵀ·J should equal first row of J (since v = [1,0])
print(f" v @ J = {cotangent @ J} (verification)")
# This is how grad() works internally!
# grad(scalar_fn)(x) is equivalent to:
# _, vjp_fn = vjp(scalar_fn, x)
# return vjp_fn(1.0) # cotangent = 1.0 for scalar output
VJP (Vector-Jacobian Product):
f(x) = [2. 7.]
vᵀ·J = [2. 1. 0.] (gradient w.r.t. x when caring only about f₁)
v @ J = [2. 1. 0.] (verification)
# Practical example: Gradient of loss through a neural network layer
def layer(params, x):
"""Single dense layer: y = tanh(Wx + b)"""
W, b = params['W'], params['b']
return jnp.tanh(W @ x + b)
def loss(params, x, y_true):
"""MSE loss"""
y_pred = layer(params, x)
return jnp.mean((y_pred - y_true)**2)
# Setup
params = {'W': jnp.array([[1., 2.], [3., 4.], [5., 6.]]),
'b': jnp.array([0.1, 0.2, 0.3])}
x = jnp.array([1.0, 1.0])
y_true = jnp.array([0.5, 0.5, 0.5])
# Using grad (internally uses VJP)
grads = grad(loss)(params, x, y_true)
print("Gradients via grad():")
print(f" dL/dW shape: {grads['W'].shape}")
print(f" dL/db shape: {grads['b'].shape}")
# Equivalent using VJP explicitly
_, vjp_fn = vjp(lambda p: loss(p, x, y_true), params)
grads_vjp, = vjp_fn(1.0) # cotangent = 1.0 for scalar loss
print(f"\nGradients via explicit VJP (should match):")
print(f" dL/dW:\n{grads_vjp['W']}")
Gradients via grad():
dL/dW shape: (3, 2)
dL/db shape: (3,)
Gradients via explicit VJP (should match):
dL/dW:
[[2.67312534e-03 2.67312534e-03]
[7.43184674e-07 7.43184674e-07]
[2.04118870e-10 2.04118870e-10]]
2.5 Hessian-Vector Products#
Computing the full Hessian is O(n²) in memory. For large n, use Hessian-vector products (HVP) instead:
This is O(n) and computes H @ v without forming H explicitly.
def hvp(f, x, v):
"""Compute Hessian-vector product H·v without forming H.
Uses forward-over-reverse mode: jvp of grad.
"""
return jvp(grad(f), (x,), (v,))[1]
# Example: Rosenbrock function (classic optimization test)
def rosenbrock(x):
"""f(x,y) = (1-x)² + 100(y-x²)²"""
return (1 - x[0])**2 + 100*(x[1] - x[0]**2)**2
x = jnp.array([1.0, 1.0]) # At the minimum
v = jnp.array([1.0, 0.0]) # Direction
# Compute HVP
hvp_result = hvp(rosenbrock, x, v)
print(f"Hessian-vector product H·v = {hvp_result}")
# Verify against full Hessian
H = hessian(rosenbrock)(x)
print(f"\nFull Hessian:\n{H}")
print(f"H @ v = {H @ v} (should match HVP)")
Hessian-vector product H·v = [ 802. -400.]
Full Hessian:
[[ 802. -400.]
[-400. 200.]]
H @ v = [ 802. -400.] (should match HVP)
# Why HVP matters: scaling to large problems
def large_quadratic(x):
"""A quadratic function in high dimensions."""
return jnp.sum(x**2) + jnp.sum(x[:-1] * x[1:]) # Banded Hessian
n = 1000
x_large = jnp.ones(n)
v_large = jnp.ones(n)
# HVP is O(n) memory and time
hvp_large = hvp(large_quadratic, x_large, v_large)
print(f"HVP for n={n}: computed in O(n) time/memory")
print(f" H·v shape: {hvp_large.shape}")
print(f" H·v[0:5] = {hvp_large[:5]}")
# Full Hessian would be O(n²) = 1,000,000 elements!
print(f"\nFull Hessian would require {n*n:,} elements ({n*n*8/1e6:.1f} MB)")
print("HVP avoids materializing the full Hessian!")
# Application: Newton-CG optimization uses HVP to solve H·δ = -g
# without ever forming H explicitly
HVP for n=1000: computed in O(n) time/memory
H·v shape: (1000,)
H·v[0:5] = [3. 4. 4. 4. 4.]
Full Hessian would require 1,000,000 elements (8.0 MB)
HVP avoids materializing the full Hessian!
3. JIT Compilation#
jit compiles functions using XLA for significant speedups.
def slow_fn(x):
"""A function with many operations."""
for _ in range(100):
x = jnp.sin(x) + jnp.cos(x)
return x.sum()
fast_fn = jit(slow_fn)
x = jnp.ones(1000)
# Warm up JIT (first call compiles)
_ = fast_fn(x)
# Time comparison
import time
start = time.time()
for _ in range(10):
_ = slow_fn(x)
slow_time = time.time() - start
start = time.time()
for _ in range(10):
_ = fast_fn(x)
fast_time = time.time() - start
print(f"Without JIT: {slow_time:.4f}s")
print(f"With JIT: {fast_time:.4f}s")
print(f"Speedup: {slow_time/fast_time:.1f}x")
Without JIT: 0.3500s
With JIT: 0.0005s
Speedup: 741.9x
Combining JIT with Grad#
Transformations compose naturally:
def loss(params):
return jnp.sum(params**2)
# Compose: JIT the gradient function
fast_grad = jit(grad(loss))
params = jnp.array([1.0, 2.0, 3.0])
print(f"Gradient: {fast_grad(params)}")
print(f"Expected: {2 * params}")
Gradient: [2. 4. 6.]
Expected: [2. 4. 6.]
4. Vectorization with vmap#
vmap automatically vectorizes functions over batch dimensions.
def single_example(x):
"""Process a single input."""
return jnp.sum(x**2)
# Without vmap: need a loop
batch = jnp.array([[1, 2, 3],
[4, 5, 6],
[7, 8, 9]], dtype=jnp.float64)
# Manual loop (slow)
results_loop = jnp.array([single_example(x) for x in batch])
# With vmap (fast, vectorized)
batched_fn = vmap(single_example)
results_vmap = batched_fn(batch)
print(f"Loop result: {results_loop}")
print(f"vmap result: {results_vmap}")
Loop result: [ 14. 77. 194.]
vmap result: [ 14. 77. 194.]
# vmap + grad: Compute gradients for a batch
def loss_fn(params, x):
return jnp.sum((params - x)**2)
# Gradient w.r.t. params for each x in batch
batched_grad = vmap(grad(loss_fn), in_axes=(None, 0))
params = jnp.array([0.0, 0.0, 0.0])
batch_x = jnp.array([[1, 2, 3],
[4, 5, 6]], dtype=jnp.float64)
grads = batched_grad(params, batch_x)
print("Gradients for each batch element:")
print(grads)
Gradients for each batch element:
[[ -2. -4. -6.]
[ -8. -10. -12.]]
5. Pytrees: Working with Nested Structures#
JAX can differentiate through nested data structures (dicts, lists, tuples).
def model(params, x):
"""Simple neural network layer."""
W = params['W']
b = params['b']
return jnp.tanh(W @ x + b)
def loss(params, x, y_true):
y_pred = model(params, x)
return jnp.sum((y_pred - y_true)**2)
# Parameters as a dictionary (pytree)
params = {
'W': jnp.array([[1.0, 2.0], [3.0, 4.0]]),
'b': jnp.array([0.1, 0.2])
}
x = jnp.array([1.0, 1.0])
y_true = jnp.array([0.5, 0.5])
# Gradient returns a pytree with same structure
grads = grad(loss)(params, x, y_true)
print("Gradients (same structure as params):")
print(f" dL/dW:\n{grads['W']}")
print(f" dL/db: {grads['b']}")
Gradients (same structure as params):
dL/dW:
[[8.01937603e-03 8.01937603e-03]
[2.22955402e-06 2.22955402e-06]]
dL/db: [8.01937603e-03 2.22955402e-06]
6. Custom Derivatives#
Sometimes you need to define custom gradient rules using custom_vjp or custom_jvp.
from jax import custom_vjp
@custom_vjp
def safe_divide(x, y):
"""Division with custom gradient handling for y=0."""
return x / y
def safe_divide_fwd(x, y):
"""Forward pass: compute output and save residuals."""
return safe_divide(x, y), (x, y)
def safe_divide_bwd(res, g):
"""Backward pass: compute gradients."""
x, y = res
# d(x/y)/dx = 1/y
# d(x/y)/dy = -x/y²
return (g / y, -g * x / y**2)
safe_divide.defvjp(safe_divide_fwd, safe_divide_bwd)
# Test
x, y = 6.0, 2.0
print(f"safe_divide({x}, {y}) = {safe_divide(x, y)}")
print(f"Gradients: {grad(safe_divide, argnums=(0,1))(x, y)}")
print(f"Expected: (1/{y}, -{x}/{y}²) = ({1/y}, {-x/y**2})")
safe_divide(6.0, 2.0) = 3.0
Gradients: (Array(0.5, dtype=float64, weak_type=True), Array(-1.5, dtype=float64, weak_type=True))
Expected: (1/2.0, -6.0/2.0²) = (0.5, -1.5)
7. Control Flow#
JAX requires special handling for control flow to maintain differentiability.
from jax import lax
# Use lax.cond instead of Python if/else
def abs_value(x):
return lax.cond(
x >= 0,
lambda: x,
lambda: -x
)
# Test
print(f"|3| = {abs_value(3.0)}")
print(f"|-3| = {abs_value(-3.0)}")
# Gradient works!
print(f"d|x|/dx at x=3: {grad(abs_value)(3.0)}")
print(f"d|x|/dx at x=-3: {grad(abs_value)(-3.0)}")
|3| = 3.0
|-3| = 3.0
d|x|/dx at x=3: 1.0
d|x|/dx at x=-3: -1.0
# Use lax.fori_loop instead of Python for loops
def power_naive(x, n):
"""Compute x^n using a loop (Python way - not JIT-able with variable n)."""
result = 1.0
for _ in range(n):
result = result * x
return result
def power_jax(x, n):
"""Compute x^n using lax.fori_loop (JIT-able)."""
def body_fn(i, result):
return result * x
return lax.fori_loop(0, n, body_fn, 1.0)
x = 2.0
n = 5
print(f"{x}^{n} = {power_jax(x, n)}")
print(f"d(x^{n})/dx = {grad(power_jax)(x, n)}")
print(f"Expected: {n}*x^{n-1} = {n * x**(n-1)}")
2.0^5 = 32.0
d(x^5)/dx = 80.0
Expected: 5*x^4 = 80.0
8. Scan: Efficient Sequential Operations#
lax.scan is like a differentiable fold/reduce operation.
def cumsum_scan(xs):
"""Cumulative sum using scan."""
def step(carry, x):
new_carry = carry + x
return new_carry, new_carry # (new_state, output)
_, cumsum = lax.scan(step, 0.0, xs)
return cumsum
xs = jnp.array([1.0, 2.0, 3.0, 4.0, 5.0])
print(f"Input: {xs}")
print(f"Cumsum (scan): {cumsum_scan(xs)}")
print(f"Cumsum (numpy): {jnp.cumsum(xs)}")
# Gradient flows through scan!
def sum_of_cumsum(xs):
return jnp.sum(cumsum_scan(xs))
print(f"\nGradient of sum(cumsum(xs)): {grad(sum_of_cumsum)(xs)}")
Input: [1. 2. 3. 4. 5.]
Cumsum (scan): [ 1. 3. 6. 10. 15.]
Cumsum (numpy): [ 1. 3. 6. 10. 15.]
Gradient of sum(cumsum(xs)): [5. 4. 3. 2. 1.]
Summary#
Function |
Purpose |
Example |
When to Use |
|---|---|---|---|
|
Scalar gradient |
|
f: Rⁿ → R (loss functions) |
|
Forward-mode Jacobian |
|
n < m (few inputs) |
|
Reverse-mode Jacobian |
|
n > m (many inputs) |
|
Jacobian-vector product |
|
Directional derivatives |
|
Vector-Jacobian product |
|
Backpropagation, custom gradients |
|
Hessian matrix |
|
Small n (full matrix needed) |
HVP |
Hessian-vector product |
|
Large n (Newton-CG) |
|
JIT compilation |
|
Speed up repeated calls |
|
Vectorization |
|
Batch processing |
|
Differentiable if/else |
|
Conditional logic |
|
Differentiable loop |
|
Sequential operations |
Key Principles:
Functions must be pure (no side effects)
Use JAX NumPy (
jnp) instead of NumPyTransformations compose:
jit(grad(vmap(f)))Use pytrees for nested parameters
Choose forward/reverse mode based on input/output dimensions