[remat3] add a remat option to custom_gradient
With remat=True, a custom_gradient function returns a rematerialization
function rem in place of its VJP function, and rem returns the output along
with the VJP function. This is the closure-style analog of
custom_vjp.defremat(fwd, rem, bwd): whatever rem closes over is saved on the
forward pass, and whatever the VJP function closes over is what rem
recomputes. rem takes the arguments again, so it need not close over them.
Under jax.remat, rem runs on the backward pass. Otherwise it runs right after
the function on the forward pass, using the VJP that defremat derives.
Running rem on the backward pass instead is what grad(jax.remat(f)) already
does, so gluing it into the forward pass keeps both behaviors available.
```python
@jax.custom_gradient(remat=True)
def sin(x):
cos_x = jnp.cos(x) # rem closes over it, so it is saved
def rem(x): # runs on the backward pass
return jnp.sin(x), lambda g: (g * cos_x,)
return jnp.sin(x), rem
```
with_logs=True is not yet supported together with remat=True.