JAX: How to accumulate jnp.array on a condition in a jit function?

Viewed 72

I want to filter a jnp.array with a condition, and accumulate to a global variable, in a jit function (so we have to use JAX control flow primitives):

import jax
import jax.numpy as jnp
from jax import jit
from jax import lax

key = jax.random.PRNGKey(42)


@jit
def get_data():
  data = jax.random.normal(key, (5, 3))
  data = data.at[-2:].set(0.)
  return data


data = get_data()
accu = data[0]


@jit
def filter(data):
  def body_fun(i):
    global accu
    accu = jnp.vstack((accu, data[i]))
    return i + 1

  lax.while_loop(lambda i: jnp.all(data[i]), body_fun, 1)

filter(data)

I expect accu.shape is (3,3) (there are three non-zero rows in data) after filter executed, but got (2,3):

Traced<ShapedArray(float32[2,3])>with<DynamicJaxprTrace(level=1/1)>

I suspect lax.while_loop iterates row 1 and 2, but global accu only got updated once, but why? Or is there any better way to accumulate jnp.array (in jit function) without using global variable?

1 Answers

Your body_fun updates accu by using a side-effect. Jax is a functional programming library: you'll want to make all the updates explicit. Meaning that accu should be in the arguments of body_fun and also returned by it after being updated.

The signature of jax.lax.while_loop is jax.lax.while_loop(cond_fun, body_fun, init_val). In your case, init_val should be the tuple (counter, accu).

The last issue is that accu should be of fixed shape. You'll want to initialize it beforehand to some shape, which in your case will be an upper bound on the shape of the final accu.

In the end, the following code works. Here, I suggest to not close over data, to make it obvious that cond_fn depends on data too. You could have a closure instead.

# Initialization
data = get_data()
accu = jax.zeros_like(data)
i = 0

def body_fn(carry): 
    i, accu, data = carry
    accu = accu.at[i].set(data[i])
    return (i + 1, accu, data)

def cond_fn(carry):
    i, accu, data = carry
    return jnp.all(data[i])

last_i, accu, _ = lax.while_loop(body_fn, (i, accu))
Related