From 82832bded27189e68ec1a27c317788e2accd981a Mon Sep 17 00:00:00 2001 From: Vaibhav Singh Date: Sat, 22 Aug 2026 16:09:57 +0000 Subject: [PATCH 1/2] MoE ring-of-experts: TC Pallas ring all-gather for the combine backward + x_sorted remat save Two default-off flags, for performance improvement on expert-parallel MoE: - moe_ring_cotangent_ag: the backward cotangent all-gather of the combine reduce-scatter runs on a TensorCore Pallas bidirectional ring kernel instead of the XLA collective, so it cannot serialize on the SparseCore collective-offload queue. Forward unchanged; backward numerically equal to lax.all_gather. - moe_x_sorted=device: remat-save the routed (expert-sorted) MoE input and its routing metadata so the backward loads them instead of re-running the dispatch all-gather and sort. Intended for small per-device batch; at larger batch the save exceeds HBM (compile-checked). New file src/maxtext/kernels/ring_ag.py: bidirectional store-and-forward ring all-gather (TC Pallas, ICI DMA) with caller-supplied collective_id and an explicit CostEstimate so the latency-hiding scheduler accounts for the DMA cost. --- src/maxtext/configs/types.py | 20 ++++ src/maxtext/kernels/ring_ag.py | 193 +++++++++++++++++++++++++++++++++ src/maxtext/layers/moe.py | 68 +++++++++++- 3 files changed, 275 insertions(+), 6 deletions(-) create mode 100644 src/maxtext/kernels/ring_ag.py diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 449a84d899..648665d073 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -884,6 +884,15 @@ class MoEGeneral(BaseModel): False, description="Whether to use Ring of Experts for sparse matmul expert parallelism.", ) + moe_ring_cotangent_ag: bool = Field( + False, + description=( + "Run the BACKWARD cotangent all-gather of the ring-of-experts combine reduce-scatter on " + "a TensorCore Pallas ring kernel instead of the XLA collective (which can serialize on " + "the SparseCore collective-offload queue). Forward is unchanged. For performance " + "improvement; numerically equal to lax.all_gather." + ), + ) moe_quantize_token_all_gather: bool = Field( False, description="Whether to quantize token activations to FP8 before All-Gather across EP shards in Ring of Experts.", @@ -1321,6 +1330,15 @@ class RematAndOffload(BaseModel): RematLocation.REMAT, description="Remat policy for the first part of a gated MoE's output.", ) + moe_x_sorted: RematLocation = Field( + RematLocation.REMAT, + description=( + "Remat policy for the routed (post-dispatch, expert-sorted) MoE input plus its small " + "routing/metadata bundle. 'device' saves them across the remat boundary so the backward " + "does not re-run the dispatch token all-gather and sort; the expert GMMs re-run from the " + "saved tensor. Default 'remat' recomputes (existing behavior)." + ), + ) moe_mlpwi_1: RematLocation = Field( RematLocation.REMAT, description="Remat policy for the second part of a gated MoE's output.", @@ -3423,6 +3441,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de "context", "mlpwi", "moe_mlpwi_0", + "moe_x_sorted", "moe_mlpwi_1", "moe_mlpwo", "mlpwi_0", @@ -4564,6 +4583,7 @@ def set_derived_values_and_validate(self) -> "RLConfig": "context", "mlpwi", "moe_mlpwi_0", + "moe_x_sorted", "moe_mlpwi_1", "moe_mlpwo", "mlpwi_0", diff --git a/src/maxtext/kernels/ring_ag.py b/src/maxtext/kernels/ring_ag.py new file mode 100644 index 0000000000..6bb6474e30 --- /dev/null +++ b/src/maxtext/kernels/ring_ag.py @@ -0,0 +1,193 @@ +"""Bidirectional store-and-forward ring all-gather (TensorCore Pallas, ICI DMA). + +Validated on v7x: tiled semantics bit-exact vs ``jax.lax.all_gather(..., tiled=True)``. Written for +IN-shard_map use: `ring_all_gather` runs the pallas_call directly on the local shard +inside an enclosing shard_map, with a caller-supplied `collective_id` (a fixed id would collide +with other in-flight Pallas collectives' barrier semaphores) and an explicit `CostEstimate` so the latency-hiding scheduler can see the ~2x-bytes +DMA cost instead of treating the custom-call as free (kernels without a cost estimate are +invisible to the scheduler's overlap decisions). + +Why a RING (not the direct-to-owner broadcast `_direct_all_gather` in moe.py): the direct pattern +sends every shard's block over multi-hop paths to all peers (bisection congestion, measured +regressing at EP=8); the ring moves each block over neighbor links only, at the validated +159-179 GB/s. +""" + +import jax +from jax import lax +from jax.experimental import pallas as pl +from jax.experimental.pallas import tpu as pltpu +import jax.numpy as jnp + +MESH = pl.DeviceIdType.MESH + + +def _strides(sizes): + """Tiled chunk strides for axes ordered outermost..innermost. strides[k]=prod(sizes[k+1:]).""" + st = [1] * len(sizes) + for i in range(len(sizes) - 2, -1, -1): + st[i] = st[i + 1] * sizes[i + 1] + return st + + +def _full_index(ref, gather_dim, start, length, split_dim=None, s_start=0, s_len=None): + """Index tuple: whole array, a [start:start+length] window on gather_dim, and (for bidi) a + [s_start:s_start+s_len] window on split_dim. The split is on a NON-gather dim so the gather_dim + offsets stay tile-aligned (splitting the gather/concat dim breaks Mosaic tile alignment).""" + idx = [pl.ds(0, ref.shape[d]) for d in range(ref.ndim)] + idx[gather_dim] = pl.ds(start, length) + if split_dim is not None: + idx[split_dim] = pl.ds(s_start, s_len) + return tuple(idx) + + +def _pick_split_dim(shape, gather_dim): + """First non-gather dim with an even extent (to halve for bidirectional). None -> force uni.""" + for d in range(len(shape)): + if d != gather_dim and shape[d] % 2 == 0: + return d + return None + + +def _neighbor(all_axes, axis, delta): + """device_id MESH dict: step `delta` along `axis` (wrap), hold every other mesh axis fixed.""" + size = lax.axis_size(axis) + nxt = lax.rem(lax.axis_index(axis) + delta + size, size) + return {a: (nxt if a == axis else lax.axis_index(a)) for a in all_axes} + + +def _ring_stage(o_ref, all_axes, axes, sizes, strides, k, chunk, gather_dim, + send_sem, recv_sem, bidi, split_dim, pipe): + """One store-and-forward ring over axes[k]; fills sizes[k] blocks of cur_len rows. + + Bidirectional splits each block along split_dim (a NON-gather dim): +dir carries the upper half, + -dir the lower half, each over the FULL gather block -> gather_dim offsets stay tile-aligned. + + Pipelining: each direction's half is further sliced into `pipe` independent sub-chunks along + split_dim, each with its OWN sem. Store-and-forward forces depth-1 PER sub-chunk (can't forward + what hasn't arrived), but the 2*pipe sub-chunks progress independently and are issued + round-robin, so at steady state 2*pipe DMAs are in flight. pipe=1 is the baseline.""" + s = sizes[k] + if s == 1: + return + cur_rows = strides[k] * chunk # gather-dim block length in rows + # Base row of this axis-k group = (sum_{j 0 else 0) * chunk + my = lax.axis_index(axes[k]) + use_bidi = bidi and split_dim is not None + + right_n = _neighbor(all_axes, axes[k], +1) + left_n = _neighbor(all_axes, axes[k], -1) + bsem = pltpu.get_barrier_semaphore() + pl.semaphore_signal(bsem, inc=1, device_id=right_n, device_id_type=MESH) + if use_bidi: + pl.semaphore_signal(bsem, inc=1, device_id=left_n, device_id_type=MESH) + pl.semaphore_wait(bsem, 2 if use_bidi else 1) + + # Pipeline "lanes": (sem_idx, neighbor, hop_sign, split_start, split_len). hop_sign=-1 forwards + # source-coord (my-i) to the +1 neighbor; +1 forwards (my+i) to the -1 neighbor. + lanes = [] + if split_dim is not None: + E = o_ref.shape[split_dim] + if use_bidi: + half = E // 2 + assert half % pipe == 0, (E, pipe) + w = half // pipe + for j in range(pipe): + lanes.append((j, right_n, -1, half + j * w, w)) # +dir, upper half + lanes.append((pipe + j, left_n, +1, j * w, w)) # -dir, lower half + else: + assert E % pipe == 0, (E, pipe) + w = E // pipe + for j in range(pipe): + lanes.append((j, right_n, -1, j * w, w)) # uni, full split sliced into pipe + else: + lanes.append((0, right_n, -1, None, None)) # uni, no split_dim: single block + + def idx(start, ss, sl): + if ss is None: + return _full_index(o_ref, gather_dim, start, cur_rows) + return _full_index(o_ref, gather_dim, start, cur_rows, split_dim=split_dim, s_start=ss, s_len=sl) + + def block_start(u): + return base_rows + lax.rem(u + s, s) * cur_rows + + prev = [None] * len(lanes) + for i in range(s - 1): + for li, (si, nb, sign, ss, sl) in enumerate(lanes): + start = block_start(my - i if sign < 0 else my + i) + if prev[li] is not None: + prev[li].wait() + prev[li] = pltpu.async_remote_copy( + o_ref.at[idx(start, ss, sl)], o_ref.at[idx(start, ss, sl)], + send_sem.at[si], recv_sem.at[si], device_id=nb, device_id_type=MESH) + for d in prev: + if d is not None: + d.wait() + + +def _kernel(w_ref, o_ref, send_sem, recv_sem, local_sem, *, + all_axes, axes, sizes, strides, chunk, gather_dim, bidi, split_dim, pipe): + # 1. place our own shard at its tiled slot: chunk index = sum coord_j * strides[j]. + my_chunk = sum(lax.axis_index(axes[j]) * strides[j] for j in range(len(axes))) + my_start = my_chunk * chunk + pltpu.make_async_copy( + w_ref.at[_full_index(w_ref, gather_dim, 0, chunk)], + o_ref.at[_full_index(o_ref, gather_dim, my_start, chunk)], + local_sem).start() + pltpu.make_async_copy( + w_ref.at[_full_index(w_ref, gather_dim, 0, chunk)], + o_ref.at[_full_index(o_ref, gather_dim, my_start, chunk)], + local_sem).wait() + # 2. nested rings, innermost axis first (largest k -> smallest stride). + for k in range(len(axes) - 1, -1, -1): + _ring_stage(o_ref, all_axes, axes, sizes, strides, k, chunk, gather_dim, + send_sem, recv_sem, bidi, split_dim, pipe) + + +def ring_all_gather(x, mesh, gather_axes, gather_dim, collective_id, *, bidi=True, pipe=1): + """In-shard_map ring all-gather of the LOCAL shard `x` over `gather_axes` (mesh axis names, + outermost..innermost), tiled on `gather_dim`. Semantics == + ``jax.lax.all_gather(x, gather_axes, axis=gather_dim, tiled=True)`` (bit-exact, pure data move). + + MUST be called INSIDE a shard_map spanning `mesh` (uses lax.axis_index on the mesh axes for MESH + device_id addressing). `collective_id` selects the barrier semaphore and must be DISTINCT from + every other concurrently in-flight Pallas collective's id.""" + if isinstance(gather_axes, str): + gather_axes = (gather_axes,) + sizes = tuple(mesh.shape[ax] for ax in gather_axes) + strides = _strides(sizes) + n_total = 1 + for s_ in sizes: + n_total *= s_ + chunk = x.shape[gather_dim] # local rows on the gather dim + out_shape = list(x.shape) + out_shape[gather_dim] = chunk * n_total + out_shape = tuple(out_shape) + all_axes = tuple(mesh.axis_names) + split_dim = _pick_split_dim(out_shape, gather_dim) if bidi else None + HBM = pltpu.MemorySpace.HBM + + def kern(w_ref, o_ref, send_sem, recv_sem, local_sem): + _kernel(w_ref, o_ref, send_sem, recv_sem, local_sem, + all_axes=all_axes, axes=tuple(gather_axes), sizes=sizes, strides=strides, + chunk=chunk, gather_dim=gather_dim, bidi=bidi, split_dim=split_dim, pipe=pipe) + + nsem = 2 * pipe # 2 directions x pipe sub-chunks (uni uses the first `pipe`) + # Cost estimate: each device receives + forwards ~2x the full gathered bytes over the ring. + # Without it the custom-call is invisible to the latency-hiding scheduler's overlap decisions. + full_bytes = 1 + for d_ in out_shape: + full_bytes *= d_ + full_bytes *= x.dtype.itemsize + return pl.pallas_call( + kern, + out_shape=jax.ShapeDtypeStruct(out_shape, x.dtype), + in_specs=[pl.BlockSpec(memory_space=HBM)], + out_specs=pl.BlockSpec(memory_space=HBM), + scratch_shapes=[pltpu.SemaphoreType.DMA((nsem,)), # send, per (dir,sub) lane + pltpu.SemaphoreType.DMA((nsem,)), # recv, per (dir,sub) lane + pltpu.SemaphoreType.DMA], # local copy + compiler_params=pltpu.CompilerParams(collective_id=collective_id), + cost_estimate=pl.CostEstimate(flops=0, bytes_accessed=2 * full_bytes, transcendentals=0), + )(x) diff --git a/src/maxtext/layers/moe.py b/src/maxtext/layers/moe.py index 1d2cf2e959..3f26802fe0 100644 --- a/src/maxtext/layers/moe.py +++ b/src/maxtext/layers/moe.py @@ -31,6 +31,7 @@ from jax.sharding import Mesh, NamedSharding from jax.sharding import PartitionSpec as P from maxtext.common import common_types as ctypes +from maxtext.kernels.ring_ag import ring_all_gather from maxtext.common.common_types import ShardMode from maxtext.kernels import megablox as mblx from maxtext.layers import attentions, linears, nnx_wrappers, quantizations @@ -57,6 +58,41 @@ COMBINE = "combine" +# Distinct barrier-semaphore ids for the in-MoE Pallas ring collectives; must not collide with +# any other collective_id in flight in the same program. +_RING_CT_AG_COLLECTIVE_ID = 55 # moe_ring_cotangent_ag: backward combine-cotangent ring all-gather + + +@functools.partial(jax.custom_vjp, nondiff_argnums=(1, 2, 3)) +def _ring_ct_reduce_scatter(output, mesh, ep_name, collective_id): + """Combine reduce-scatter whose BACKWARD cotangent all-gather runs on the TC ring kernel. + + FORWARD: byte-identical to the stock path -- the plain + ``jax.lax.psum_scatter(output, ep_name, scatter_dimension=0, tiled=True)`` (also what a remat + recompute re-traces: the primal is the plain collective, so no Pallas DMA ever runs inside a + rematted region). BACKWARD: the autodiff transpose of the tiled psum_scatter is a tiled EP + all-gather of the combine cotangent; as an XLA collective it can ride the SparseCore + collective-offload queue and serialize behind other SC work. Here it runs on the bidirectional + store-and-forward TensorCore ring kernel instead (ICI DMAs issued from the TC, where the + backward has slack), numerically == ``lax.all_gather`` (pure tiled data move, bit-exact). + """ + return jax.lax.psum_scatter(output, ep_name, scatter_dimension=0, tiled=True) + + +def _ring_ct_rs_fwd(output, mesh, ep_name, collective_id): + return _ring_ct_reduce_scatter(output, mesh, ep_name, collective_id), None + + +def _ring_ct_rs_bwd(mesh, ep_name, collective_id, _res, ct): + return (ring_all_gather(ct, mesh, (ep_name,), 0, collective_id),) + + +_ring_ct_reduce_scatter.defvjp(_ring_ct_rs_fwd, _ring_ct_rs_bwd) + + + + + @struct.dataclass class RouteMetadata: """EP communication state needed to undo the forward all-to-all after expert computation.""" @@ -2358,6 +2394,18 @@ def _moe_body( x, routing, route_metadata = route( x, logits, pre_bias_logits, rngs, input_ids=sharded_input_ids ) + # moe_x_sorted: tag the routed expert input and its small routing/metadata bundle for + # the remat policy. With moe_x_sorted=device the backward LOADS these instead of + # re-running route() -- removing the rematted dispatch token all-gather and sort from the + # backward (the up-projection weight gradient needs the sorted input anyway). The + # routing/metadata leaves (indices, group sizes, weights -- tiny) must be saved too, else + # the sort re-runs just to reproduce them. Tags are inert under the default + # moe_x_sorted=remat. + _cn = lambda t: adc.checkpoint_name(t, "moe_x_sorted") if isinstance(t, jax.Array) else t + # tree.map so a QArray x (moe_quantize_token_all_gather) gets its qvalue/scale leaves tagged + x = jax.tree.map(_cn, x) + routing = jax.tree.map(_cn, routing) + route_metadata = jax.tree.map(_cn, route_metadata) if self.config.mlp_bias: w0_bias, w1_bias, wo_bias = self.transform_bias( @@ -2411,12 +2459,20 @@ def _moe_body( self.moe_expert_input_dim // self.get_tensor_parallelism_size(), ), ) - output = jax.lax.psum_scatter( - output, - self._expert_parallelism_name, - scatter_dimension=0, - tiled=True, - ) + if ( + getattr(self.config, "moe_ring_cotangent_ag", False) + and isinstance(self._expert_parallelism_name, str) + ): + output = _ring_ct_reduce_scatter( + output, self.mesh, self._expert_parallelism_name, _RING_CT_AG_COLLECTIVE_ID + ) + else: + output = jax.lax.psum_scatter( + output, + self._expert_parallelism_name, + scatter_dimension=0, + tiled=True, + ) return output, routing.lb_loss, routing.bias_updates if self.get_expert_parallelism_size() > 1: From 63cdc7c6d36eb4efbf5f5676ff99e4ef7a3a5465 Mon Sep 17 00:00:00 2001 From: Vaibhav Singh Date: Sat, 22 Aug 2026 18:50:32 +0000 Subject: [PATCH 2/2] review: single descriptor for the local shard copy (start/wait on one object) --- src/maxtext/kernels/ring_ag.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/src/maxtext/kernels/ring_ag.py b/src/maxtext/kernels/ring_ag.py index 6bb6474e30..40a6bcc4f0 100644 --- a/src/maxtext/kernels/ring_ag.py +++ b/src/maxtext/kernels/ring_ag.py @@ -131,14 +131,12 @@ def _kernel(w_ref, o_ref, send_sem, recv_sem, local_sem, *, # 1. place our own shard at its tiled slot: chunk index = sum coord_j * strides[j]. my_chunk = sum(lax.axis_index(axes[j]) * strides[j] for j in range(len(axes))) my_start = my_chunk * chunk - pltpu.make_async_copy( + local_copy = pltpu.make_async_copy( w_ref.at[_full_index(w_ref, gather_dim, 0, chunk)], o_ref.at[_full_index(o_ref, gather_dim, my_start, chunk)], - local_sem).start() - pltpu.make_async_copy( - w_ref.at[_full_index(w_ref, gather_dim, 0, chunk)], - o_ref.at[_full_index(o_ref, gather_dim, my_start, chunk)], - local_sem).wait() + local_sem) + local_copy.start() + local_copy.wait() # 2. nested rings, innermost axis first (largest k -> smallest stride). for k in range(len(axes) - 1, -1, -1): _ring_stage(o_ref, all_axes, axes, sizes, strides, k, chunk, gather_dim,