Why is the return type of jax.grad different to other jax transformations in the following scenario?
Consider a function to be transformed by JAX which takes a custom container as an argument
import jax
import jax.numpy as jnp
from collections import namedtuple
# Define NTContainer using namedtuple factory
NTContainer = namedtuple(
'NTContainer',['name','flag','array']
)
# Instantiate container
container = NTContainer(
name='composite-container',
flag=0,
array=jnp.reshape(jnp.arange(6,dtype='f4'),(3,2)),
)
# Test function
def test_func(container : NTContainer):
if container.flag == 0:
return container.array.sum()
return container.array.prod()
JAX needs to be informed how to handle NTContainer correctly (default namedtuple pytree cannot be used)
# Register NTContainer pytree to handle static and traceable members
def unpack_NTC(c):
active, passive = (c.array,), (c.name, c.flag)
return active, passive
def repack_NTC(passive, active):
(name, flag), (array,) = passive, active
return NTContainer(name,flag,value)
jax.tree_util.register_pytree_node(
NTContainer, unpack_NTC, repack_NTC,
)
Now performing several jax transforms and calling with container results in
jax.jit(test_func)(container)
# DeviceArray(15., dtype=float32)
jax.vmap(test_func)(container)
# DeviceArray([1., 5., 9.], dtype=float32)
jax.grad(test_func)(container)
# NTContainer(name='composite-container',
# flag=0,
# array=DeviceArray([[1., 1.],
# [1., 1.],
# [1., 1.]], dtype=float32))
Why does the jax.grad transform call return a NTContainer rather than a DeviceArray?