Skip to content

Chapter 9

Overview

JAX requires pure functions, so standard Python if/while won't work in jitted code. Use JAX control flow primitives.

Topics

1. lax.cond - Conditional Execution

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

def absolute_value(x):
 return lax.cond(
 x < 0,
 lambda: -x, # True branch
 lambda: x # False branch
)

print(absolute_value(-5.0)) # 5.0
print(absolute_value(5.0)) # 5.0

2. lax.fori_loop - Loops with Bounds

def sum_n(n):
 def body(i, carry):
 return carry + i

 return lax.fori_loop(0, n, body, 0)

print(sum_n(5)) # 0+1+2+3+4 = 10

3. lax.while_loop - Dynamic Loops

def while_count(x):
 def cond(state):
 return state < 10

 def body(state):
 return state + 1

 return lax.while_loop(cond, body, x)

print(while_count(0)) # 10

4. lax.scan - Recurrent Processing

def scan_sum(xs):
 def body(carry, x):
 new_carry = carry + x
 return new_carry, new_carry

 final, ys = lax.scan(body, 0, xs)
 return ys

print(scan_sum(jnp.array([1, 2, 3, 4, 5])))
# Output

With JIT

@jax.jit
def jitted_conditional(x):
 return lax.cond(x > 0, lambda: x**2, lambda: -x**2)

print(jitted_conditional(3.0)) # 9.0
print(jitted_conditional(-3.0)) # -9.0

Summary

  • Use lax.cond for if statements
  • Use lax.fori_loop for bounded loops
  • Use lax.while_loop for dynamic loops
  • Use lax.scan for recurrence
  • All work with jit and grad!