I'm trying to write a (matrix) exponentiate-by-squaring algorithm in JAX. Unfortunately, I don't understand traced variables very well, which is complicating matters.
My code is:
import numpy as np
import jax
import jax.numpy as jnp
import jax.lax as jlax
@partial(jax.jit, static_argnums=(1,))
def matpow(A, n):
dim = A.shape[0]
return jlax.switch(
n,
[lambda: jnp.identity(dim),
lambda: A,
lambda: jlax.cond(
jnp.floor_divide(n, 2) == jnp.true_divide(n, 2),
lambda: matpow(jnp.dot(A, A), jnp.floor_divide(n, 2)),
lambda: jnp.dot(A, matpow(jnp.dot(A, A), jnp.floor_divide(n, 2)))
)])
However, attempting to run this with, say, matpow(2 * jnp.eye(4), 5) throws an error midway through compilation:
ValueError: Non-hashable static arguments are not supported, as this can lead to unexpected cache-misses. Static argument (index 1) of type <class 'jax.interpreters.partial_eval.DynamicJaxprTracer'> for function matpow is non-hashable.
... I have no idea what that means, to be quite honest, but it's doubly confusing because as far as I can tell n should just be an integer and therefore has a trivial hash.
Other attempts have included: using jnp.binary_repr (not yet implemented), using np.binary_repr (TracerIntegerConversionError, even though n is marked static), and a 'wrapped' version of the above recursive function in which I defined a separate function recpow internally to matpow (hit the recursion limit.)
What do I need to do to make this code work?