[remat3] use eval_jaxpr_p in RematTraced.expand to reduce trace times
```python
# Trace-time benchmark for RematTraced.expand via eval_jaxpr_p.
# Run with: JAX_REMAT3=1 python bench_remat3_expand.py
import time
import jax
import jax.numpy as jnp
from jax._src import core
from jax._src import ad_checkpoint
from jax._src.interpreters import partial_eval as pe
assert jax.config.jax_remat3
new_expand = ad_checkpoint.RematTraced.expand
def old_expand(self, *args): # the previous implementation, for comparison
return core.jaxpr_as_fun(self.jaxpr)(*args)
def make_block(n, nested):
inner = jax.checkpoint(lambda x: jnp.cos(jnp.cos(x)))
def f(x):
for _ in range(n):
x = jnp.sin(x)
return inner(x) if nested else x
return jax.checkpoint(f)
def measure_lower_pass(n, calls=1, nested=False, k=5):
# time the hi->lo lowering pass on a fresh traced jaxpr each iteration
ts = []
for _ in range(k):
block = make_block(n, nested)
def g(x):
for _ in range(calls):
x = block(x)
return x
hi = jax.jit(g).trace(1.0).jaxpr
t0 = time.perf_counter()
pe.lower_jaxpr2(hi.jaxpr)
ts.append(time.perf_counter() - t0)
return min(ts)
def measure_full_lower(n, calls, k=3):
# time full jit lowering, same remat block called `calls` times
ts = []
for _ in range(k):
jax.clear_caches()
block = make_block(n, nested=False)
def g(x):
for _ in range(calls):
x = block(x)
return x
t0 = time.perf_counter()
jax.jit(g).lower(1.0)
ts.append(time.perf_counter() - t0)
return min(ts)
for name, expand in [('before', old_expand), ('after', new_expand)]:
ad_checkpoint.RematTraced.expand = expand
print(f'--- {name} ---')
for n in [100, 400, 1600]:
t = measure_lower_pass(n)
print(f' hi->lo pass, lo body, {n:5d} eqns: {t*1e3:8.2f} ms')
for n in [400, 1600]:
t = measure_lower_pass(n, nested=True)
print(f' hi->lo pass, hi body (nested), {n:5d} eqns: {t*1e3:6.2f} ms')
t = measure_lower_pass(400, calls=8, nested=True)
print(f' hi->lo pass, hi body, 400 eqns x 8 calls: {t*1e3:8.2f} ms')
for calls in [1, 8]:
t = measure_full_lower(400, calls)
print(f' full jit lower, 400-eqn block x {calls} call(s): {t*1e3:6.2f} ms')
```
On my laptop:
```
--- before ---
hi->lo pass, lo body, 100 eqns: 0.58 ms
hi->lo pass, lo body, 400 eqns: 2.22 ms
hi->lo pass, lo body, 1600 eqns: 8.87 ms
hi->lo pass, hi body (nested), 400 eqns: 2.29 ms
hi->lo pass, hi body (nested), 1600 eqns: 9.47 ms
hi->lo pass, hi body, 400 eqns x 8 calls: 18.87 ms
full jit lower, 400-eqn block x 1 call(s): 8.33 ms
full jit lower, 400-eqn block x 8 call(s): 37.50 ms
--- after ---
hi->lo pass, lo body, 100 eqns: 0.04 ms
hi->lo pass, lo body, 400 eqns: 0.04 ms
hi->lo pass, lo body, 1600 eqns: 0.05 ms
hi->lo pass, hi body (nested), 400 eqns: 0.27 ms
hi->lo pass, hi body (nested), 1600 eqns: 0.83 ms
hi->lo pass, hi body, 400 eqns x 8 calls: 0.36 ms
full jit lower, 400-eqn block x 1 call(s): 5.70 ms
full jit lower, 400-eqn block x 8 call(s): 6.06 ms
```