Skip to content

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