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!