With the single replica broadcast feature enabled, we are able to see performance gain over restoration. However, if broadcast_memory_limit_bytes is set, the trace and compilation overhead become significant which almost kill the point of broadcast. As of today, if you check the broadcast implementation _src/multihost/multislice.py:
- L292–307 _merge_globalized_replicas: jax.jit(lambda tree: ...) builds a new lambda on every call. JAX caches by function identity, so every call traces, lowers and compiles again, even when the input signature is identical.
- L343–382 the chunk loop in broadcast_one_replica_to_all: one compile per chunk, and jax.block_until_ready(out_subtree) at L381 makes them strictly serial (compile c1, run c1, compile c2, …).
It would be nice if we could hide the compilation time behind the reading, since the shape of the PyTree is known ahead of reading. At the moment, this is blocking because _src/serialization/jax_array_handlers.py:
- L1739 deserialized = await _deserialize_arrays(...): the whole primary-replica read finishes before anything else happens.
- L1781 multislice.broadcast_one_replica_to_all(...): only then do the chunk compiles start.
- Every chunk's input signature (args[i].global_shape, dtype, single_replica_shardings, mesh) is known at L1720, before the read starts.
- L1760 create_zeros (hosts outside the primary replica): same pattern, a jit defined inside the function. This one is off the critical path.
You could reproduce this by set jax.config.update("jax_log_compiles", True), or JAX_LOG_COMPILES=1:
"""multislice.broadcast_one_replica_to_all recompiles its merge program for
every chunk, on every call, on the critical path. CPU, single process,
8 fake devices, mesh (replica=2, model=4)."""
import os
os.environ["XLA_FLAGS"] = "--xla_force_host_platform_device_count=8"
import logging, time
import jax, jax.numpy as jnp, numpy as np
from jax._src import monitoring
from orbax.checkpoint._src.multihost import multislice
jax.config.update("jax_log_compiles", True)
compile_s = []
monitoring.register_event_duration_secs_listener(
lambda e, d, **kw: e == "/jax/core/compile/backend_compile_duration" and compile_s.append(d))
devices = np.array(jax.devices()).reshape(2, 4)
global_mesh = jax.sharding.Mesh(devices, ("replica", "model"))
replica0 = jax.sharding.Mesh(devices[:1], ("replica", "model"))
P = jax.sharding.PartitionSpec
def make_tree(num_layers=16):
"""Like a transformer checkpoint: many leaves, identical per-layer shapes."""
leaves = []
for _ in range(num_layers):
leaves.append(jax.device_put(jnp.ones((512, 512)), jax.sharding.NamedSharding(replica0, P(None, "model"))))
leaves.append(jax.device_put(jnp.ones((512,)), jax.sharding.NamedSharding(replica0, P("model"))))
return tuple(leaves)
per_layer_bytes = 512 * 128 * 4 + 128 * 4 # per-device bytes of one (w, b) pair
for call in range(2): # identical inputs both times
tree = make_tree()
jax.block_until_ready(tree)
compile_s.clear()
t0 = time.time()
_, n = multislice.broadcast_one_replica_to_all(
tree, global_mesh, replica_axis_index=0, is_source=True,
memory_limit_bytes=2 * per_layer_bytes)
total = time.time() - t0
logging.warning(f"call {call}: {n} broadcasts, {len(compile_s)} XLA compiles, "
f"compile {sum(compile_s):.2f}s of {total:.2f}s total ({sum(compile_s)/total:.0%})")
Output from my run:
call 0: 8 broadcasts, 28 XLA compiles, compile 0.35s of 0.48s total (73%)
call 1: 8 broadcasts, 8 XLA compiles, compile 0.11s of 0.20s total (54%)
With jax_log_compiles on, you'll see 16 Compiling jit() lines with byte-identical signatures (float32[2,512,512], float32[2,512], ..., same shardings), 8 per call:
- Call 0's extra 20 compiles are one-time eager warm-up for the zeros placeholders and expand_dims; they're cached afterwards.
- Call 1 recompiles all 8 chunks even though every chunk has an identical signature and the inputs match call 0. With a stable function, it would compile once in call 0 and not at all in call 1.
Once implemented, this should also help fix issue#3181
With the single replica broadcast feature enabled, we are able to see performance gain over restoration. However, if
broadcast_memory_limit_bytesis set, the trace and compilation overhead become significant which almost kill the point of broadcast. As of today, if you check the broadcast implementation _src/multihost/multislice.py:It would be nice if we could hide the compilation time behind the reading, since the shape of the PyTree is known ahead of reading. At the moment, this is blocking because _src/serialization/jax_array_handlers.py:
You could reproduce this by set jax.config.update("jax_log_compiles", True), or JAX_LOG_COMPILES=1:
Output from my run:
call 0: 8 broadcasts, 28 XLA compiles, compile 0.35s of 0.48s total (73%)
call 1: 8 broadcasts, 8 XLA compiles, compile 0.11s of 0.20s total (54%)
With jax_log_compiles on, you'll see 16 Compiling jit() lines with byte-identical signatures (float32[2,512,512], float32[2,512], ..., same shardings), 8 per call:
Once implemented, this should also help fix issue#3181