How to improve this toy Jax optimizer code with while loops and saved history?

Viewed 164

I'm writing a custom optimizer I want JIT-able with Jax which features 1) breaking on maximum steps reached 2) breaking on a tolerance reached, and 3) saving the history of the steps taken. I'm relatively new to some of this stuff in Jax, but reading the docs I have this solution:

import jax, jax.numpy as jnp

@jax.jit
def optimizer(x, tol = 1, max_steps = 5):
    
    def cond(arg):
        step, x, history = arg
        return (step < max_steps) & (x > tol)
    
    def body(arg):
        step, x, history = arg
        x = x / 2 # simulate taking an optimizer step
        history = history.at[step].set(x) # simulate saving current step
        return (step + 1, x, history)

    return jax.lax.while_loop(
        cond,
        body,
        (0, x, jnp.full(max_steps, jnp.nan))
    )

optimizer(10.) # works

My question is whether this can be improved in some way? In particular, is there a way to avoid pre-allocating the history? This isn't ideal since the real thing is alot more complicated than a single array and there's obviously the potential for wasted memory if tolerance is reached well before the maximum steps.

2 Answers

is there a way to avoid pre-allocating the history?

No, as I understand JAX

in JAX, 'type' includes shape, that is, the in and out data shape of the body function MUST be the same, otherwise, say dynamic grow history use jnp.vstack((history, x)), JAX will consider it as side effect.

There is a way, if you think the tolerance will be often reached before the maximum number of steps.

JAX implements sparse matrices (and pytrees of them), in jax.experimental.sparse. They will have the same shape as the maximum history size, and therefore satisfy the "fixed size" requirement for XLA, but of course, will only store nonzero elements in memory.

Related