Skip to content

broadcast_one_replica_to_all recompiles per chunk on the critical path: compile during the checkpoint read #3606

Description

@lyglst

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, …).
Image

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

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions