Skip to content

Chapter 5

Overview

jax.jit compiles Python functions to machine code via XLA, providing significant speedups.


What is JIT Compilation?

JIT = Just-In-Time - compile when first called, then run fast.

Traditional Python is Slow

def slow_sum(n):
 """Pure Python loop"""
 total = 0
 for i in range(n):
 total += i
 return total

# Slow
import time
start = time.time()
result = slow_sum(1000000)
print(f"Time: {time.time() - start:.4f}s") # ~100ms

JAX JIT Compiles to Machine Code

import jax
import jax.numpy as jnp
import time

def fast_sum(n):
 """JAX can trace and compile this"""
 return jnp.sum(jnp.arange(n))

# Compile with JIT
fast_sum_jit = jax.jit(fast_sum)

# First call
start = time.time()
result = fast_sum_jit(1000000)
print(f"Compile time: {time.time() - start:.4f}s")

# Subsequent calls
start = time.time()
result = fast_sum_jit(1000000)
print(f"Execution time: {time.time() - start:.6f}s")

Result: 100x-1000x faster after compilation!

-

How JIT Works: The Tracing Process

Step 1: Write Pure Function

import jax.numpy as jnp

def f(x):
 return x**2 + 2*x + 1

Step 2: JAX Traces Function

import jax

# When you call jax.jit(f):
# 1. JAX creates "abstract" symbolic input
# 2. Traces through f, recording operations
# 3. Builds computation graph
# 4. Compiles graph with XLA

f_jit = jax.jit(f)

Step 3: Compilation

Traced operations:
 - x² (Square)
 - 2*x (Multiply)
 - + (Add)
 - + (Add)
 ↓
Optimization:
 - Fuse operations
 - Eliminate redundancy
 - Generate machine code
 ↓
Compiled code (native machine code)
 ↓
GPU/TPU kernels (if available)

Step 4: Execution

# First call
result1 = f_jit(3.0)

# Later calls
result2 = f_jit(3.0)
result3 = f_jit(3.0)

Basic Usage

Simple JIT

import jax
import jax.numpy as jnp

# Define function
def f(x):
 return jnp.sin(x) + jnp.cos(x)

# Create JIT version
f_jit = jax.jit(f)

# Use it
result = f_jit(1.0) # First call: compile
result = f_jit(1.0) # Subsequent calls: fast!

# Or use decorator
@jax.jit
def g(x):
 return x**2 + 1

result = g(3.0)

Multiple Arguments

import jax
import jax.numpy as jnp

@jax.jit
def dot_product(x, y):
 return jnp.dot(x, y)

x = jnp.array([1., 2., 3.])
y = jnp.array([4., 5., 6.])

result = dot_product(x, y) # 32.0

Nested JIT

import jax
import jax.numpy as jnp

@jax.jit
def inner(x):
 return x**2

@jax.jit
def outer(x):
 return inner(x) + 1

# Works fine - JIT composition
result = outer(3.0) # 10.0

Static Arguments: static_argnums

Sometimes you don't want to recompile for different argument values.

Problem: Recompilation

import jax
import jax.numpy as jnp

@jax.jit
def f(x, n):
 return jnp.power(x, n) # x^n

# Different n → different computation graph
result1 = f(2.0, 3) # Compiles for n=3
result2 = f(2.0, 4) # Recompiles for n=4 (slow!)

Solution: Static Arguments

import jax
import jax.numpy as jnp

@jax.jit
def f(x, n):
 # n is not traced - it's part of compilation decision
 return jnp.power(x, n)

# Compiler sees n as "static" (compile-time constant)
# Different n values trigger recompilation
result1 = f(2.0, 3) # Compiles for n=3
result2 = f(2.0, 3) # Reuses same compilation
result3 = f(2.0, 4) # Recompiles for n=4

Shape Polymorphism

Problem: Different Shapes

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
 return jnp.sum(x) # Sum doesn't depend on shape

# Different shapes
result1 = f(jnp.array([1., 2., 3.])) # Shape (3,)
result2 = f(jnp.array([1., 2., 3., 4.])) # Shape (4,) - recompiles!

Solution: Polymorphic Shapes

import jax
import jax.numpy as jnp

# JAX can compile once for variable shapes
def f(x):
 return jnp.sum(x)

# With proper design, one compilation handles multiple shapes
# (JAX tries to be smart about this)

f_jit = jax.jit(f)
result1 = f_jit(jnp.array([1., 2., 3.]))
result2 = f_jit(jnp.array([1., 2., 3., 4.])) # Recompiles (different shape)

Static vs Dynamic Shapes

Static Shapes (JIT-friendly)

import jax
import jax.numpy as jnp

@jax.jit
def f(x, y):
 return x + y

x = jnp.array([1., 2., 3.]) # Shape (3,) known at compile
y = jnp.array([4., 5., 6.]) # Shape (3,) known at compile

result = f(x, y) # Compiles for (3,) + (3,)

Dynamic Shapes (Complex)

import jax
import jax.numpy as jnp
import jax.lax as lax

@jax.jit
def f(x, n):
 # This is tricky - shape depends on n (runtime value)
 return lax.fori_loop(0, n, lambda i, carry: carry + x, 0.0)

# Works but requires careful design
result = f(jnp.array([1., 2., 3.]), 5)

Combining JIT with Other Transformations

JIT + Grad

import jax
import jax.numpy as jnp
from jax import grad, jit

def loss(params, x, y):
 pred = params @ x
 return jnp.mean((pred - y)**2)

# Slow
grad_loss = grad(loss)

# Fast
grad_loss_jit = jit(grad(loss))

params = jnp.array([1., 2., 3.])
x = jnp.array([1., 2., 3.])
y = jnp.array([5., 6., 7.])

grads = grad_loss_jit(params, x, y) # First: compile, then execute
grads = grad_loss_jit(params, x, y) # Subsequent: fast!

JIT + Vmap

import jax
import jax.numpy as jnp
from jax import vmap, jit

def f(x):
 return x**2 + 1

# Vectorize then JIT
f_jit_vmap = jit(vmap(f))

batch = jnp.array([1., 2., 3., 4., 5.])
result = f_jit_vmap(batch) # Fast and vectorized!

JIT + Grad + Vmap

import jax
import jax.numpy as jnp
from jax import grad, vmap, jit

def loss(w, x, y):
 return jnp.mean((w * x - y)**2)

# Gradients for batch, compiled
grad_loss_batch = jit(vmap(grad(loss, argnums=0)))

# Use on batches
W_batch = jnp.array([[1., 2.], [3., 4.]]) # (2, 2)
X_batch = jnp.array([[1., 2.], [3., 4.]]) # (2, 2)
y_batch = jnp.array([1., 2.]) # (2,)

grads = grad_loss_batch(W_batch, X_batch, y_batch)

Common Issues

Issue 1: Data-Dependent Shapes

import jax
import jax.numpy as jnp

@jax.jit
def bad_fn(x):
 # Shape-dependent operation!
 if x[0] > 0:
 return x[:5] # Conditional shape
 else:
 return x[:10]

# This won't work predictably

Solution: Use jax.lax.cond for conditional shapes.

Issue 2: Python Loops

import jax
import jax.numpy as jnp

@jax.jit
def bad_fn(x, n):
 result = 0
 for i in range(n): # Python loop (can't trace)
 result += x[i]
 return result

# Use jax.lax.fori_loop instead
import jax.lax as lax

@jax.jit
def good_fn(x, n):
 def body(i, result):
 return result + x[i]
 return lax.fori_loop(0, n, body, 0)

Issue 3: Print Statements

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
 print(x) # Only prints during tracing!
 return x**2

# First call
result = f(3.0)

# Subsequent calls
result = f(3.0)

Solution: Use jax.debug.print for debugging JIT code.


Debugging JIT

Use debug.print

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
 jax.debug.print("x = {}", x) # Prints during execution
 return x**2

result = f(3.0)

Disable JIT for Debugging

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
 return x**2

# Disable JIT for debugging
with jax.disable_jit():
 result = f(3.0) # Runs without compilation

Find Tracer Issues

import jax
import jax.numpy as jnp

@jax.jit
def f(x):
 # If x is Tracer during execution, you'll see errors
 # This indicates JIT incompatibility
 y = x > 0 # This might cause issues
 return y

# Check what type x is during tracing

Performance Profiling

Measure Compilation Time

import jax
import jax.numpy as jnp
import time

@jax.jit
def f(x):
 return jnp.sum(jnp.sin(x)**2)

x = jnp.arange(10000.0)

# Compile
start = time.time()
result = f(x)
compile_time = time.time() - start
print(f"Compile time: {compile_time:.3f}s")

# Execution
start = time.time()
result = f(x)
exec_time = time.time() - start
print(f"Execution time: {exec_time:.6f}s")

Profile with JAX

import jax
import jax.numpy as jnp
from jax import profiler

@jax.jit
def f(x):
 return jnp.sum(x**2)

x = jnp.arange(10000.0)

# Profile
profiler.trace(lambda: f(x))

Best Practices

1. JIT at Module Level

# Good
@jax.jit
def train_step(params, batch):
 # Training logic
 return new_params

# Use in loop
for batch in batches:
 params = train_step(params, batch)
# Bad
for batch in batches:
 @jax.jit # Creates new JIT version each iteration!
 def train_step(params):
 return new_params

 params = train_step(params)

2. Match Shapes

# Good
X_batch = jnp.array([[1, 2], [3, 4]]) # (2, 2) each iteration
for X in X_batch:
 result = f_jit(X) # Always (2,)

# Bad
batches = [jnp.array([1]), jnp.array([1, 2, 3])]
for batch in batches:
 result = f_jit(batch) # Recompiles for each shape!

3. Combine with grad

# Efficient
@jax.jit
def loss_and_grad(params, x, y):
 loss = compute_loss(params, x, y)
 grad_params = jax.grad(compute_loss)(params, x, y)
 return loss, grad_params

# Or simpler
train_step = jax.jit(
 lambda p, x, y: (
 compute_loss(p, x, y),
 jax.grad(compute_loss)(p, x, y)
)
)

-

Summary

JIT COMPILATION CHECKLIST:

Use @jax.jit decorator for functions called repeatedly
Ensure all dependencies are pure (no side effects)
Keep shapes consistent (matching shapes = fewer recompiles)
Use static_argnums for non-array arguments
Combine with grad/vmap for powerful optimization
Disable with jax.disable_jit() for debugging
Use jax.debug.print for debugging
Avoid data-dependent shapes and Python control flow

-

Next Steps

Checkpoint: Time a non-JIT vs JIT version of a simple function

-

Last Updated: 2026-08-09