[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)
```