Skip to content

Chapter 9: Control Flow in JAX

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: [1, 3, 6, 10, 15]

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!