jax
ae1ef100 - [remat3] use eval_jaxpr_p in RematTraced.expand to reduce trace times

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