Chapter 3: Pure Functions & Immutability¶
Overview¶
JAX requires functional programming. This chapter covers pure functions and immutable data structures essential for JAX.
What is a Pure Function?¶
A pure function satisfies: 1. Deterministic: Same input → always same output 2. No side effects: Doesn't modify external state 3. No I/O: Doesn't read/write files or network 4. Referentially transparent: Can be replaced with return value
Example: Pure vs Impure¶
# ❌ IMPURE: Global state
total = 0
def add_impure(x):
global total
total += x
print(f"Total: {total}") # I/O side effect
return total
result1 = add_impure(5)
result2 = add_impure(5)
# Same input, different outputs (5 vs 10)!
# ✅ PURE: No side effects, deterministic
def add_pure(x):
return x + x
result1 = add_pure(5) # 10
result2 = add_pure(5) # 10 (always same)
Why Pure Functions Matter in JAX¶
The Tracing Problem¶
When you call jax.grad(f), JAX traces your function:
import jax
import jax.numpy as jnp
def f_pure(x):
return x**2 + jnp.sin(x)
# JAX traces through f with symbolic values:
# 1. Executes with symbolic x
# 2. Records operations: x², sin(x), +
# 3. Computes derivative of recorded graph
grad_f = jax.grad(f_pure)
print(grad_f(3.0)) # Works: 2*3 + cos(3) ≈ 6.96
With side effects, tracing breaks:
import jax
import jax.numpy as jnp
global_counter = 0
def f_impure(x):
global global_counter
global_counter += 1 # Side effect: modifies external state
return x**2
# During tracing:
# JAX increments global_counter, but doesn't record it
# When computing gradient, counter state is inconsistent
grad_f = jax.grad(f_impure)
print(grad_f(3.0)) # Unpredictable result!
Composition Requires Purity¶
JAX transformations compose only when functions are pure:
import jax
import jax.numpy as jnp
# Pure function
def f(x):
return x**2
# Compose transformations
gradient_of_vmap = jax.grad(jax.vmap(f))
# = Gradient of (Apply f to batch)
# = Works correctly for pure functions
# With impure function:
global_state = []
def g_impure(x):
global_state.append(x) # Side effect
return x**2
# Composing transformations on impure functions:
try:
jax.grad(jax.vmap(g_impure)) # Unpredictable behavior!
except:
pass
Eliminating Side Effects¶
Pattern 1: Return Results Instead of Modifying¶
# ❌ IMPURE: Modifies input list
def append_impure(lst, item):
lst.append(item) # Modifies original!
return lst
my_list = [1, 2, 3]
result = append_impure(my_list, 4)
print(my_list) # [1, 2, 3, 4] - MODIFIED!
# ✅ PURE: Returns new list
def append_pure(lst, item):
return lst + [item] # New list
my_list = [1, 2, 3]
result = append_pure(my_list, 4)
print(my_list) # [1, 2, 3] - unchanged
print(result) # [1, 2, 3, 4]
Pattern 2: Pass Dependencies as Arguments¶
# ❌ IMPURE: Uses global
config = {"learning_rate": 0.01}
def train_step_impure(weights):
lr = config["learning_rate"] # Depends on global state
return weights - lr * gradients
# ✅ PURE: Takes all dependencies as arguments
def train_step_pure(weights, learning_rate):
return weights - learning_rate * gradients
# Now reusable with different learning rates
weights1 = train_step_pure(weights, 0.01)
weights2 = train_step_pure(weights, 0.001)
Pattern 3: Return Multiple Values¶
# ❌ IMPURE: Modifies objects
class State:
def __init__(self):
self.count = 0
self.values = []
state = State()
def process_impure(x):
state.count += 1
state.values.append(x)
return x**2
# ✅ PURE: Returns new state
def process_pure(state, x):
new_count = state['count'] + 1
new_values = state['values'] + [x]
new_state = {'count': new_count, 'values': new_values}
result = x**2
return result, new_state
state = {'count': 0, 'values': []}
result, state = process_pure(state, 5)
Immutability in JAX¶
JAX Arrays are Immutable¶
import jax.numpy as jnp
x = jnp.array([1, 2, 3])
# ❌ Can't modify in place
# x[0] = 999 # TypeError: 'Array' object does not support item assignment
# ✅ Create new array with .at[]
x = x.at[0].set(999)
print(x) # [999 2 3]
Immutable Update Methods¶
import jax.numpy as jnp
x = jnp.array([1, 2, 3, 4, 5])
# set: Replace value
x = x.at[0].set(999)
print(x) # [999 2 3 4 5]
# add: Add to value
x = jnp.array([1, 2, 3])
x = x.at[0].add(10)
print(x) # [11 2 3]
# multiply: Multiply value
x = jnp.array([1, 2, 3])
x = x.at[1].multiply(100)
print(x) # [1 200 3]
# subtract, divide, min, max
x = x.at[2].subtract(5) # x[2] -= 5
x = x.at[2].divide(2) # x[2] /= 2
# Batch updates
x = jnp.array([1, 2, 3, 4, 5])
x = x.at[0, 2, 4](/0,-2,-4/).set(999)
print(x) # [999 2 999 4 999]
Chaining Updates¶
import jax.numpy as jnp
x = jnp.array([1, 2, 3, 4, 5])
# Chain multiple updates
x = (x.at[0].add(10)
.at[2].multiply(100)
.at[4].set(999))
print(x) # [11 2 300 4 999]
Deep Immutability¶
import jax.numpy as jnp
# Nested structures
data = {
'A': jnp.array([1, 2, 3]),
'B': jnp.array([4, 5, 6])
}
# ✅ Update array in dict
data['A'] = data['A'].at[0].set(999)
# Alternative: create new dict
data = {
**data,
'A': data['A'].at[0].set(999)
}
Functional Programming Patterns¶
Higher-Order Functions¶
Functions that return functions:
# Return a configured function
def make_adder(n):
def add(x):
return x + n
return add
add_5 = make_adder(5)
print(add_5(10)) # 15
print(add_5(20)) # 25
Closures¶
Outer function's variables captured:
def make_multiplier(factor):
def multiply(x):
return x * factor
return multiply
times_3 = make_multiplier(3)
times_5 = make_multiplier(5)
print(times_3(10)) # 30
print(times_5(10)) # 50
Function Composition¶
Combining functions:
def compose(f, g):
"""Returns h(x) = f(g(x))"""
def h(x):
return f(g(x))
return h
# Example
def double(x):
return x * 2
def add_one(x):
return x + 1
double_then_add = compose(add_one, double)
print(double_then_add(5)) # add_one(double(5)) = 11
Handling State Functionally¶
State as Arguments/Returns¶
import jax.numpy as jnp
# Traditional: state in object
# class Accumulator:
# def __init__(self):
# self.total = 0
# def add(self, x):
# self.total += x # Impure!
# Functional: state as argument/return
def accumulate_step(state, x):
"""Returns (result, new_state)"""
new_total = state['total'] + x
new_state = {**state, 'total': new_total}
return new_total, new_state
# Usage
state = {'total': 0}
result1, state = accumulate_step(state, 5)
result2, state = accumulate_step(state, 3)
result3, state = accumulate_step(state, 7)
print(state) # {'total': 15}
Training Loop Example¶
import jax
import jax.numpy as jnp
from jax import grad
def loss(params, x, y):
pred = params['w'] * x + params['b']
return jnp.mean((pred - y)**2)
def train_step(params, x, y, learning_rate):
"""Pure: takes state, returns new state"""
grads = grad(loss)(params, x, y)
new_params = {
'w': params['w'] - learning_rate * grads['w'],
'b': params['b'] - learning_rate * grads['b']
}
loss_value = loss(new_params, x, y)
return new_params, loss_value
# Usage
params = {'w': 0.0, 'b': 0.0}
x, y = jnp.array([1., 2., 3.]), jnp.array([2., 4., 6.])
for step in range(100):
params, loss_val = train_step(params, x, y, 0.01)
print(params) # Learned parameters
JAX-Specific Patterns¶
Using lax for Control Flow¶
import jax
import jax.lax as lax
import jax.numpy as jnp
# Pure conditional
def absolute_value(x):
return lax.cond(
x < 0,
lambda: -x, # True branch
lambda: x # False branch
)
print(absolute_value(-5)) # 5
print(absolute_value(5)) # 5
# Pure loop (scan)
def sum_n(n):
def body(carry, x):
return carry + x, None
result, _ = lax.scan(body, 0, jnp.arange(n))
return result
print(sum_n(10)) # 45 (sum of 0 to 9)
Using JAX Pytrees¶
JAX can work with nested structures automatically:
import jax
import jax.numpy as jnp
# Nested structure (pytree)
tree = {
'a': jnp.array([1., 2., 3.]),
'b': jnp.array([4., 5.]),
'c': [jnp.array([6.]), jnp.array([7., 8.])]
}
# JAX operations work on trees
def add_one_to_all(tree):
# JAX automatically maps over structure
return jax.tree_map(lambda x: x + 1, tree)
result = add_one_to_all(tree)
print(result['a']) # [2. 3. 4.]
# Flatten tree for operations
leaves, treedef = jax.tree_util.tree_flatten(tree)
print(len(leaves)) # 4 (flattened arrays)
# Reconstruct
tree_reconstructed = jax.tree_util.tree_unflatten(treedef, leaves)
Common Mistakes¶
Mistake 1: Assuming Mutability¶
import jax.numpy as jnp
x = jnp.array([1, 2, 3])
# ❌ This doesn't modify x
y = x.reshape((1, 3))
# x is still shape (3,)
# ❌ This doesn't modify x
z = x + 5
# x is still [1, 2, 3]
# ✅ Reassign to modify "state"
x = x.reshape((1, 3))
x = x + 5
Mistake 2: Modifying Function Arguments¶
# ❌ IMPURE: Modifies argument
def process_list_bad(lst):
lst.append(999) # Modifies input!
return sum(lst)
# ✅ PURE: Returns new list
def process_list_good(lst):
new_lst = lst + [999] # New list
return sum(new_lst)
Mistake 3: Using Global Variables¶
# ❌ IMPURE
global_config = {'lr': 0.01}
def train_step(params):
return params - global_config['lr'] * grads
# ✅ PURE
def train_step(params, lr):
return params - lr * grads
Best Practices¶
1. Design for Immutability¶
# Prefer functional approach
def process_data(data, config):
# Don't modify data or config
processed = transform(data)
new_config = {**config, 'updated': True}
return processed, new_config
2. Use Return Values for State¶
# Return modified state
state, loss = train_step(state, batch)
state, new_loss = train_step(state, next_batch)
3. Pass Dependencies as Arguments¶
# All needed information as arguments
def compute(x, y, z, config):
# Don't rely on globals
return f(x, y, z, config['param'])
4. Use JAX Utilities for Control Flow¶
# Use lax.cond, lax.scan, lax.while_loop
# Not Python if/for statements
Summary¶
PURE FUNCTION CHECKLIST:
✓ Same input → same output (deterministic)
✓ No global state modifications
✓ No I/O (files, network, print in grad)
✓ No randomness (use jax.random properly)
✓ Returns all output values
✓ Takes all inputs as arguments
Next Steps¶
- 01 Automatic Differentiation Jax.Grad - Use pure functions for gradients
- 03 Vectorization Jax.Vmap - Compose with other transformations
Checkpoint: Explain why JAX requires pure functions (hint: tracing and composition)
Last Updated: 2026-08-09