jax
8ee4e8eb - [remat3] add a remat option to custom_gradient

Commit
13 days ago
[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.
Author
Committer
Parents
Loading