Chapter 11: Building Neural Networks with JAX¶
Overview¶
JAX is functional, so neural networks are just functions mapping parameters to outputs.
Topics¶
1. Manual Network Implementation¶
import jax
import jax.numpy as jnp
from jax import random, grad
def relu(x):
return jnp.maximum(0, x)
def forward(params, x):
w1, b1, w2, b2 = params
h = relu(jnp.dot(x, w1) + b1)
return jnp.dot(h, w2) + b2
def loss(params, x, y):
pred = forward(params, x)
return jnp.mean((pred - y)**2)
# Initialize
key = random.PRNGKey(0)
w1 = random.normal(key, (10, 5)) * 0.1
b1 = jnp.zeros(5)
w2 = random.normal(key, (5, 1)) * 0.1
b2 = jnp.zeros(1)
params = (w1, b1, w2, b2)
# Gradient descent
grad_loss = grad(loss)
for step in range(100):
grads = grad_loss(params, x_train, y_train)
# Update params (immutable way)
params = jax.tree_map(lambda p, g: p - 0.01 * g, params, grads)
2. Initializing Layers¶
import jax.numpy as jnp
from jax import random
def init_params(key, input_size, hidden_size, output_size):
key1, key2, key3, key4 = random.split(key, 4)
params = {
'w1': random.normal(key1, (input_size, hidden_size)) * 0.1,
'b1': jnp.zeros(hidden_size),
'w2': random.normal(key2, (hidden_size, output_size)) * 0.1,
'b2': jnp.zeros(output_size),
}
return params
3. Training Loop¶
def train(params, x, y, learning_rate, epochs):
grad_loss = grad(loss)
for epoch in range(epochs):
grads = grad_loss(params, x, y)
# Update all parameters
params = {
k: params[k] - learning_rate * grads[k]
for k in params.keys()
}
if epoch % 10 == 0:
print(f"Epoch {epoch}, Loss: {loss(params, x, y):.4f}")
return params
Summary¶
- Define params as pytrees (dicts, tuples)
- Use jax.tree_map for updates
- grad() works with nested params
- vmap for batch processing
- jit for compilation