jax
8131b0bc - [remat3] replace custom_remat with custom_vjp.defremat

Commit
16 days ago
[remat3] replace custom_remat with custom_vjp.defremat Under jax_remat3 and jax_custom_vjp3, f.defremat(fwd, rem, bwd) customizes how a custom_vjp function is rematerialized when differentiated under jax.remat, in place of the default of rematerializing its defvjp fwd rule. It reuses the helper custom_vjp that the default remat rule already builds, so the separate CustomRemat primitive and jax.custom_remat are removed. Unless defvjp is also called, the defremat rules also define the VJP used outside of jax.remat, by running rem right after fwd on the forward pass. Running rem on the backward pass instead is what grad(jax.remat(f)) already does, so gluing rem into the forward pass keeps both behaviors available: grad(f) does not rematerialize, and grad(jax.remat(f)) uses the custom rules. Before: ```python sin = jax.custom_remat( jnp.sin, lambda policy, x: (jnp.sin(x), jnp.cos(x)), # f_fwd: save cos(x) lambda cos_x, x: (jnp.sin(x), cos_x), # f_rem: reuse it lambda cos_x, g: (cos_x * g,)) # f_bwd ``` After: ```python @jax.custom_vjp def sin(x): return jnp.sin(x) def sin_fwd(x): return jnp.sin(x), jnp.cos(x) # save cos(x) def sin_rem(cos_x, x): return jnp.sin(x), cos_x # no need to recompute cos(x) def sin_bwd(cos_x, g): return (cos_x * g,) sin.defremat(sin_fwd, sin_rem, sin_bwd) ```
Author
Committer
Parents
Loading