From cdf4d2979e9febec0602c945ad6bdad9f253df44 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 22:28:04 +0000 Subject: [PATCH 01/19] Mirror the bytes a recomputed CP attention keeps for backward A size-only mirror of the executor's record_for_backward path, beside the code it mirrors: per stage, the Q/K/V flex consumes (padded copies, else contiguous copies of permuted views, or kept gathers and fetch buffers), flex's output and LSE at the execution length, logical copies when padded, and a merge-tape clone for every producing stage after the first. It reproduces the traced CP2 ranks: 0.971 GB for an aligned single stage and 3.07 GB for a padded local plus full-query remote stage. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/executor.py | 80 ++++++++++++ .../test_context_parallel_retained_bytes.py | 118 ++++++++++++++++++ 2 files changed, 198 insertions(+) create mode 100644 tests/unit/test_context_parallel_retained_bytes.py diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 4015011fe..78a65069f 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -35,6 +35,7 @@ DkvReducePlan, ExactMaskMetadata, FlexMaskSpec, + RankRuntimePlan, StageExecutionSpec, StagePlan, TokenRange, @@ -1550,6 +1551,85 @@ def _merge_stage_output_grads_from_tape( return stage_out_grads, stage_lse_grads +def retained_stage_record_bytes( + rank_plan: RankRuntimePlan, + *, + q_heads: int, + kv_heads: int, + head_dim: int, + value_head_dim: int, + element_size: int, + block_size: SparseBlockSize, +) -> int: + """Bytes a recomputed attention keeps for backward beyond its own-row tensors. + + A size-only mirror of ``_forward_stage_records`` with ``record_for_backward`` + and of ``_run_stage_attention``; keep the three in step. Every stage that + runs keeps the Q/K/V flex consumed: a padded copy when the execution length + differs, else a contiguous copy of the permuted ``q_flat``/``k_flat`` view; + partial-range gathers and remote fetch buffers are kept as the stage's + inputs. It keeps flex's output and LSE at the execution length, and + logical-length copies of them when padded. Every producing stage after the + first keeps a merge-tape clone of the accumulators. Accumulators themselves + are transient. + """ + own = int(rank_plan.local_valid_lengths[0]) if rank_plan.local_valid_lengths else 0 + accum_size = 4 if element_size < 4 else element_size + q_row = q_heads * head_dim * element_size + k_row = kv_heads * head_dim * element_size + v_row = kv_heads * value_head_dim * element_size + out_row = q_heads * value_head_dim * element_size + lse_row = q_heads * 4 + tape_row = q_heads * (value_head_dim + 1) * accum_size + total = 0 + local_produced = False + tapes: list[int] = [] + for stage in _ordered_stage_plans(rank_plan.stage_plans): + if not (stage.q_len > 0 and stage.k_len > 0 and stage.slices): + continue + q_len = _logical_stage_q_len(stage) + k_len = _logical_stage_k_len(stage) + q_pad, k_pad, _family = select_sparse_execution_family( + is_local_stage=bool(stage.is_local_stage), + q_len=int(stage.q_len), + k_len=int(stage.k_len), + block_size=block_size, + ) + q_full = _ranges_cover_full_length(stage.owner_local_q_ranges, length=own) + # Queries: the full range is a permuted view of q_flat; a partial range + # is gathered into a contiguous tensor kept as the stage's input. + if not q_full: + total += q_row * q_len + if q_pad != q_len: + total += q_row * q_pad + elif q_full and q_heads > 1: + total += q_row * q_len + # Keys and values: local ranges as for queries; remote ones land in + # contiguous head-major fetch buffers kept as the stage's inputs. + k_full = bool(stage.is_local_stage) and _ranges_cover_full_length( + stage.owner_local_k_ranges, length=own + ) + if not k_full: + total += (k_row + v_row) * k_len + if k_pad != k_len: + total += (k_row + v_row) * k_pad + elif k_full and kv_heads > 1: + total += (k_row + v_row) * k_len + total += (out_row + lse_row) * q_pad + if q_pad != q_len: + total += (out_row + lse_row) * q_len + tape = tape_row * (own if q_full else q_len) + if stage.is_local_stage: + local_produced = True + else: + tapes.append(tape) + # The first producing stage records no tape: the local stage when it ran, + # else whichever remote stage is ready first, so drop the smallest tape. + if tapes and not local_produced: + tapes.remove(min(tapes)) + return total + sum(tapes) + + def _forward_stage_records( *, q_flat: torch.Tensor, diff --git a/tests/unit/test_context_parallel_retained_bytes.py b/tests/unit/test_context_parallel_retained_bytes.py new file mode 100644 index 000000000..fd5e0e643 --- /dev/null +++ b/tests/unit/test_context_parallel_retained_bytes.py @@ -0,0 +1,118 @@ +"""Size-only retained-set mirror of the CP attention executor's recompute path.""" + +import pytest + +pytest.importorskip("triton") + +from art.megatron.context_parallel.executor import ( # noqa: E402 + retained_stage_record_bytes, +) +from art.megatron.context_parallel.types import ( # noqa: E402 + RankRuntimePlan, + StagePlan, + TokenRange, +) + +# Qwen3.6-35B-A3B attention: 16 query heads, 2 KV heads of 256, BF16, and the +# H200 flash block for a 256-wide head (128 query, 64 key rows). +GEOMETRY = dict( + q_heads=16, + kv_heads=2, + head_dim=256, + value_head_dim=256, + element_size=2, + block_size=(128, 64), +) +Q, KV, OUT, TAPE = 8192, 2048, 8192 + 64, 16 * 257 * 4 + + +def stage(index, *, local, q, k, q_len=None, k_len=None, own=None, source=0): + q_ranges = ( + (TokenRange(0, q),) if own is None or q == own else (TokenRange(1, 1 + q),) + ) + return StagePlan( + stage_index=index, + source_rank=source, + is_local_stage=local, + slices=("slice",), # ty: ignore[invalid-argument-type] + owner_local_q_ranges=q_ranges if q else (), + owner_local_k_ranges=(TokenRange(0, k),) if k else (), + q_len=q if q_len is None else q_len, + k_len=k if k_len is None else k_len, + ) + + +def plan(own, *stages): + return RankRuntimePlan( + rank=0, + original_seq_len=own, + token_layout_index=None, # ty: ignore[invalid-argument-type] + local_valid_lengths=(own,), + local_row_ranges=(TokenRange(0, own),), + stage_plans=stages, + remote_dkv_reduce_plan=None, # ty: ignore[invalid-argument-type] + ) + + +def test_aligned_single_stage_copies_views_without_output_copies(): + # Real-data rank 0: one aligned local stage. Contiguous copies of the + # permuted Q/K/V views, flex output and LSE; no padding, copies or tape. + rows = 52480 + retained = retained_stage_record_bytes( + plan(rows, stage(0, local=True, q=rows, k=rows)), **GEOMETRY + ) + assert retained == Q * rows + KV * rows + OUT * rows # 0.971 GB traced + + +def test_unaligned_single_stage_pads_and_copies_the_logical_output(): + rows = 52481 + q_pad, k_pad = 411 * 128, 821 * 64 + retained = retained_stage_record_bytes( + plan(rows, stage(0, local=True, q=rows, k=rows)), **GEOMETRY + ) + assert retained == Q * q_pad + KV * k_pad + OUT * q_pad + OUT * rows + + +def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): + # Real-data rank 1: a local stage and a remote stage over all of its + # queries. Planner lengths round up, so both stages pad. + own, remote_k = 44314, 16504 + local = stage(0, local=True, q=own, k=own, q_len=44352, k_len=44352) + remote = stage(1, local=False, q=own, k=remote_k, q_len=44352, k_len=16512) + retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) + q_pad = 347 * 128 # 44,416, as flex's traced output size shows + local_bytes = Q * q_pad + KV * 44352 + OUT * q_pad + OUT * own + remote_bytes = ( + Q * q_pad + KV * remote_k + KV * 16512 + OUT * q_pad + OUT * own + TAPE * own + ) + assert retained == local_bytes + remote_bytes + assert 3.05e9 < retained < 3.10e9 # 3.07 GB traced + + +def test_partial_query_remote_stage_keeps_its_gather_and_a_partial_tape(): + # Random-data rank 1: a large local stage and a 768-row remote stage. + own = 105153 + local = stage(0, local=True, q=own, k=own, q_len=105216, k_len=105216) + remote = stage(1, local=False, q=768, k=20608, own=own) + retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) + local_bytes = Q * 105216 + KV * 105216 + OUT * 105216 + OUT * own + # Aligned: the query gather and fetch buffers feed flex without copies. + remote_bytes = Q * 768 + KV * 20608 + OUT * 768 + TAPE * 768 + assert retained == local_bytes + remote_bytes + + +def test_empty_remote_stage_and_missing_local_stage(): + rows = 1024 + empty = stage(1, local=False, q=0, k=0) + alone = retained_stage_record_bytes( + plan(rows, stage(0, local=True, q=rows, k=rows), empty), **GEOMETRY + ) + assert alone == Q * rows + KV * rows + OUT * rows + # Without a local stage, the first ready remote stage records no tape; the + # mirror drops the smallest so it never under-counts the order. + small = stage(1, local=False, q=256, k=256, own=rows) + full = stage(2, local=False, q=rows, k=512) + both = retained_stage_record_bytes(plan(rows, small, full), **GEOMETRY) + without_small_tape = retained_stage_record_bytes(plan(rows, full), **GEOMETRY) + assert both - without_small_tape == Q * 256 + KV * 256 + OUT * 256 + TAPE * rows + assert retained_stage_record_bytes(plan(rows), **GEOMETRY) == 0 From 5899a2dee8ce9b16013bf2afdf0827ec445f91b6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 22:47:10 +0000 Subject: [PATCH 02/19] Price CP2 recompute per rank on each layer type's layout The checkpoint floor priced every rank on the busiest rank's rows with one mixer width, which over-counted boundaries and GDN and under-counted a rank whose attention runs two full-query stages. At CP2/TP1 with ART's CP core attention, price each rank on its own layouts instead: saved boundaries by each layer input's layout, a recomputed attention layer on attention rows plus what the executor keeps for backward, a GDN layer on GDN rows, and the MoE stage on that layer's rows with routed rows on the EP share. Admission takes the largest rank. Elsewhere the busiest-rank floor is unchanged. A rank can receive more routed rows than it holds, so routed rows are no longer clamped to local rows. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/runtime.py | 49 ++++ src/art/trainer_rank/_impl.py | 215 +++++++++++++++++- .../test_trainer_rank_admission_inputs.py | 1 + .../test_trainer_rank_checkpoint_memory.py | 9 +- tests/unit/test_trainer_rank_layout_memory.py | 161 +++++++++++++ tests/unit/test_trainer_rank_moe_memory.py | 2 +- 6 files changed, 420 insertions(+), 17 deletions(-) create mode 100644 tests/unit/test_trainer_rank_layout_memory.py diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 531cb8829..2f7ac2dad 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -433,6 +433,55 @@ def context_parallel_rank_model_token_counts( ) +def context_parallel_rank_layouts( + *, + group_ids: torch.Tensor, + parent_ids: torch.Tensor, + topology: ParallelTopology, + config: ContextParallelConfig, + original_seq_len: int, + build_gdn_execution_spec: bool, + gdn_planner_config: Any | None = None, +) -> tuple[tuple[int, ...], tuple[int, ...] | None, tuple[RankRuntimePlan, ...]]: + """Each CP rank's attention rows, GDN rows and attention stage plan. + + Uses the cached planning bundle and per-rank runtime plans that execution + builds, so a memory estimate sees the layouts the ranks will run. + """ + planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = ( + _get_or_build_planning_bundle( + group_ids=group_ids, + parent_ids=parent_ids, + topology=topology, + config=config, + original_seq_len=original_seq_len, + build_gdn_execution_spec=build_gdn_execution_spec, + ) + ) + attention = tuple(bundle.token_layout_index.token_counts_by_rank) + gdn = None + if build_gdn_execution_spec: + gdn = tuple( + _plan_gdn_global_execution( + planning_key=planning_key, + bundle=bundle, + topology=topology, + gdn_planner_config=gdn_planner_config, + ).gdn_token_counts_by_rank + ) + plans = tuple( + _get_or_build_bundle_rank_plan( + planning_key=planning_key, + bundle=bundle, + original_seq_len=original_seq_len, + target_rank=rank, + block_size=config.block_size, + ) + for rank in range(len(attention)) + ) + return attention, gdn, plans + + def context_parallel_model_token_total( *, group_ids: torch.Tensor, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 8dfc41239..389a865c7 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1031,6 +1031,17 @@ def signature(self) -> _MemorySignature: _AnyForwardPlan = _FlatForwardPlan | _SplitForwardPlan +@dataclass(frozen=True) +class _GroupLayout: + """One packed group's CP layouts on every rank, for layout-aware pricing.""" + + attention_rows: tuple[int, ...] + gdn_rows: tuple[int, ...] | None + # What each rank's recomputed CP attention keeps for backward beyond its + # own-row activations (``retained_stage_record_bytes``). + attention_retained: tuple[int, ...] + + @dataclass(frozen=True) class _SubforwardCost: """Memory terms of one candidate subforward while all graphs stay live. @@ -4157,8 +4168,10 @@ def _moe_workspace_bytes( raise ValueError("Invalid constructor converted-weight stages") if type(shared) is not int or not 0 <= shared <= coefficient: raise ValueError("Invalid constructor shared-expert coefficient") - routed = rows if routed_rows is None else max(0, min(rows, routed_rows)) - return (rows - routed) * shared + ( + # A rank can receive more routed rows than it holds (an uneven CP + # split); the shared part then stays on the routed rows' count. + routed = rows if routed_rows is None else max(0, routed_rows) + return max(0, rows - routed) * shared + ( max( routed * coefficient, *(routed * per_row + fixed for per_row, fixed in stages), @@ -4172,11 +4185,14 @@ def _checkpoint_memory_floor( group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None = None, routed_rows: tuple[int, ...] | None = None, + layouts: tuple[_GroupLayout, ...] | None = None, ) -> tuple[int, int]: """Conservative saved-boundary charge and one recomputed layer's workspace. ``routed_rows`` are each group's balanced dispatched rows per rank - (``_plan_group_routed_rows``); by default, its local rows. + (``_plan_group_routed_rows``); by default, its local rows. With + ``layouts`` (``_plan_group_layouts``), price every rank on its own CP + layouts instead of the busiest rank's rows (``_layout_checkpoint_floor``). Count actual local full/uniform/1 boundaries, including aliases, rather than claiming measured distinct storage. Only this call's new groups enter the term; already-live graphs remain in the availability baseline. @@ -4241,9 +4257,12 @@ def _checkpoint_memory_floor( or getattr(decoder, "_forward_pre_hooks", None) ): return 0, 0 + refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + routed = (None,) * len(group_rows) if routed_rows is None else routed_rows + if layouts is not None and all(grad for _, grad in group_rows): + return self._layout_checkpoint_floor(decoder.layers, refs, routed, layouts) retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 - refs = (None,) * len(group_rows) if slot_refs is None else slot_refs # Beside the mixer, the recomputed layer keeps its post-mixer residual # and pre-MLP norm output, and its MoE stage its routing state. mixer = ( @@ -4256,7 +4275,6 @@ def _checkpoint_memory_floor( if gradient_rows else 0 ) - routed = (None,) * len(group_rows) if routed_rows is None else routed_rows workspace = max( self._moe_workspace_bytes( rows, routed_rows=dispatched, checkpoint_grad=grad, slot_ref=ref @@ -4270,6 +4288,164 @@ def _checkpoint_memory_floor( workspace += self._te_workspace_growth_bytes() return retained, workspace + def _layout_checkpoint_floor( + self, + layers: Sequence[torch.nn.Module], + refs: tuple["LoRASlotRef | None", ...], + routed: tuple[int | None, ...], + layouts: tuple[_GroupLayout, ...], + ) -> tuple[int, int]: + """Each CP rank's boundaries and recomputed layer on its own layouts. + + A saved layer input arrives in the GDN layout when the layer follows a + GDN layer in its island (``_art_gdn_island_boundary``), and in the + attention layout otherwise. A recomputed attention layer keeps its + activations on its attention rows plus what the CP executor keeps for + backward; a GDN layer keeps the GDN width on its GDN rows. Either one's + residual, norm, routing state and MoE stage use that layer's rows, and + routed rows the EP share. Traced CP2 ranks: boundaries match exactly and + attention within 0.3%. Returns the largest rank's boundaries and the + rest of the largest rank total, so the two sum to that total. + """ + hidden = self._hidden_size * 2 + gdn_inputs = sum( + getattr( + getattr(layer, "_art_gdn_island_boundary", None), "input_layout", "" + ) + == "gdn" + for layer in layers + ) + attention_inputs = len(layers) - gdn_inputs + widths = self._recomputed_mixer_widths(stage_buffers=False) + moe = self._checkpoint_moe_bytes_per_token() + beside = 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe else 0 + retained_by_rank: list[int] = [] + totals: list[int] = [] + for rank in range(len(layouts[0].attention_rows)): + retained = workspace = 0 + for ref, dispatched, layout in zip(refs, routed, layouts, strict=True): + attention = max(1, layout.attention_rows[rank]) + gdn = ( + attention + if layout.gdn_rows is None + else max(1, layout.gdn_rows[rank]) + ) + retained += hidden * (attention_inputs * attention + gdn_inputs * gdn) + for kind, rows, extra in ( + ("attention", attention, layout.attention_retained[rank]), + ("gdn", gdn, 0), + ): + if kind not in widths: + continue + stage = ( + rows * (widths[kind] + beside) + + extra + + self._moe_workspace_bytes( + rows, + routed_rows=dispatched, + checkpoint_grad=True, + slot_ref=ref, + ) + ) + workspace = max(workspace, stage) + retained_by_rank.append(retained) + totals.append(retained + workspace) + retained = max(retained_by_rank) + workspace = max(totals) - retained + if moe: + workspace += self._te_workspace_growth_bytes() + return retained, workspace + + def _plan_group_layouts( + self, plan: _FlatForwardPlan + ) -> tuple[_GroupLayout, ...] | None: + """Every rank's CP layouts per group, where layout pricing is modeled. + + Only CP2 at TP1/PP1 with gradient groups, ART's CP core attention with + no softmax offset, and GDN layers marked with island boundaries; the + executor's retained set is validated there. Elsewhere ``None`` keeps + the busiest-rank pricing. + """ + _dp, tp, cp, pp = plan.signature.topology + if (tp, cp, pp) != (1, 2, 1) or not plan.groups: + return None + if not all(group.grad_enabled for group in plan.groups): + return None + geometry = self._geometry + if not geometry.num_attention_heads or not geometry.kv_channels: + return None + try: + decoder = _language_model(self.runtime.model[0]).decoder + from art.megatron.context_parallel.core_attention import ( + ArtContextParallelCoreAttention, + ) + except (AttributeError, RuntimeError, ModuleNotFoundError): + return None + for layer in decoder.layers: + boundary = getattr(layer, "_art_gdn_island_boundary", None) + if boundary is not None and boundary.is_gdn: + continue + if self._gdn_layers and boundary is None: + return None + core = getattr( + getattr(layer, "self_attention", None), "core_attention", None + ) + if ( + type(core) is not ArtContextParallelCoreAttention + or getattr(core, "softmax_offset", None) is not None + ): + return None + from art.megatron.context_parallel.executor import retained_stage_record_bytes + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.flex_attn.compiled import flash_sparse_block_size_for_head_dim + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + _gdn_planner_config_for_provider, + ) + + topology = self._topology() + handler = self.runtime.model_support_handler + config = _context_parallel_config_for_provider( + self.runtime.provider, self.device, handler + ) + head = int(geometry.kv_channels) + block = flash_sparse_block_size_for_head_dim( + head_dim=head, head_dim_v=head, device=self.device + ) + layouts = [] + for group in plan.groups: + batch = _pad_packed_batch(group.packed, multiple=int(topology.tp)) + attention, gdn, rank_plans = context_parallel_rank_layouts( + group_ids=batch.group_ids, + parent_ids=batch.parent_ids, + topology=topology, + config=config, + original_seq_len=int(batch.tokens.shape[1]), + build_gdn_execution_spec=handler.build_gdn_execution_spec, + gdn_planner_config=_gdn_planner_config_for_provider( + self.runtime.provider, handler + ), + ) + layouts.append( + _GroupLayout( + attention_rows=attention, + gdn_rows=gdn, + attention_retained=tuple( + retained_stage_record_bytes( + rank_plan, + q_heads=int(geometry.num_attention_heads), + kv_heads=int(geometry.num_query_groups), + head_dim=head, + value_head_dim=head, + element_size=self._param_dtype_size, + block_size=block, + ) + for rank_plan in rank_plans + ), + ) + ) + return tuple(layouts) + def _te_workspace_growth_bytes(self) -> int: """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. @@ -4361,6 +4537,7 @@ def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: gdn_segments=plan.grad_segment_count, group_rows=self._plan_group_rows(plan), group_routed_rows=self._plan_group_routed_rows(plan), + group_layouts=self._plan_group_layouts(plan), slot_refs=tuple(g.slot_ref for g in plan.groups), head_workspace_bytes=self._plan_head_workspace_bytes(plan), checkpoint_floor=_gdn_memory.plan_floor(self, plan), @@ -4378,6 +4555,7 @@ def _subforward_cost( gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), group_routed_rows: tuple[int, ...] | None = None, + group_layouts: tuple[_GroupLayout, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, checkpoint_floor: tuple[int, int] = (0, 0), @@ -4392,6 +4570,7 @@ def _subforward_cost( gdn_segments=gdn_segments, group_rows=group_rows, group_routed_rows=group_routed_rows, + group_layouts=group_layouts, slot_refs=slot_refs, head_workspace_bytes=head_workspace_bytes, checkpoint_floor=checkpoint_floor, @@ -4399,7 +4578,7 @@ def _subforward_cost( include_checkpoint_input_gradient=False, ) checkpoint_retained, checkpoint_workspace = self._checkpoint_memory_floor( - group_rows, slot_refs, group_routed_rows + group_rows, slot_refs, group_routed_rows, group_layouts ) retained = self._retained_memory_bytes( signature, @@ -7221,6 +7400,7 @@ def _memory_check( gdn_segments=forward.grad_segment_count, group_rows=self._plan_group_rows(forward), group_routed_rows=self._plan_group_routed_rows(forward), + group_layouts=self._plan_group_layouts(forward), slot_refs=tuple(g.slot_ref for g in forward.groups), head_workspace_bytes=self._plan_head_workspace_bytes(forward), checkpoint_floor=_gdn_memory.plan_floor(self, forward), @@ -8068,6 +8248,7 @@ def _estimate_required_memory_bytes_from_values( gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), group_routed_rows: tuple[int, ...] | None = None, + group_layouts: tuple[_GroupLayout, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, checkpoint_floor: tuple[int, int] = (0, 0), @@ -8160,7 +8341,7 @@ def _estimate_required_memory_bytes_from_values( ), ) retained, workspace = self._checkpoint_memory_floor( - group_rows, slot_refs, group_routed_rows + group_rows, slot_refs, group_routed_rows, group_layouts ) static_compute = max( static_compute, @@ -8299,18 +8480,26 @@ def _recomputed_mixer_bytes_per_token(self) -> int: more rows; its hidden-width input exchange and value-width output allowance price that (88 KB measured at CP2, 94 KB priced). """ + return max(self._recomputed_mixer_widths().values(), default=0) + + def _recomputed_mixer_widths(self, *, stage_buffers: bool = True) -> dict[str, int]: + """Recomputed attention and GDN mixer bytes per row, by layer type. + + ``stage_buffers=False`` leaves out the CP attention stage allowance, for + callers that price the executor's stage buffers from its stage plan. + """ geometry = self._geometry hidden = self._hidden_size tp = max(1, self._topology_key()[1]) cp = self._topology_key()[2] > 1 - widths = [] + widths: dict[str, float] = {} if self._gdn_layers < self._num_layers: attention, _gdn = self._mixer_activation_widths() - if cp: + if cp and stage_buffers: q = geometry.num_attention_heads * geometry.kv_channels or hidden kv = geometry.num_query_groups * geometry.kv_channels or hidden attention += (3 * q + 2 * kv) / tp - widths.append(attention) + widths["attention"] = attention if self._gdn_layers: key = geometry.gdn_key_heads * geometry.gdn_key_head_dim value = geometry.gdn_value_heads * geometry.gdn_value_head_dim @@ -8319,8 +8508,10 @@ def _recomputed_mixer_bytes_per_token(self) -> int: gdn = hidden + (2 * key + normalized + 6 * value + chunk) / tp if cp: gdn += hidden + value / tp - widths.append(gdn) - return int(max(widths, default=0) * self._param_dtype_size) + widths["gdn"] = gdn + return { + kind: int(width * self._param_dtype_size) for kind, width in widths.items() + } def _gdn_segment_layer_bytes(self) -> float: """Initial and final fp32 recurrent states plus convolution history.""" diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index 3481bfd6d..66b4bf812 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -33,6 +33,7 @@ def assert_plan_values(rank, plan, values): assert values["gdn_segments"] == plan.grad_segment_count assert values["group_rows"] == rank._plan_group_rows(plan) assert values["group_routed_rows"] == rank._plan_group_routed_rows(plan) + assert values["group_layouts"] == rank._plan_group_layouts(plan) assert values["head_workspace_bytes"] == rank._plan_head_workspace_bytes(plan) assert values["checkpoint_floor"] == _gdn_memory.plan_floor(rank, plan) assert values["retained_tokens"] == rank._plan_retained_tokens(plan) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 1d1087ab6..10ef298d1 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -534,10 +534,11 @@ def test_routed_rows_move_only_the_routed_moe_part(): fewer = r._checkpoint_memory_floor(((local, True),), None, (routed,)) assert fewer[0] == retained assert workspace - fewer[1] == (local - routed) * (188416 - 8192) - # Never more routed rows than local ones. - assert r._moe_workspace_bytes( - 10, routed_rows=20, checkpoint_grad=True - ) == r._moe_workspace_bytes(10, checkpoint_grad=True) + # A rank can receive more routed rows than it holds: all are priced, and + # the shared part stays on the routed count. + assert r._moe_workspace_bytes(10, routed_rows=20, checkpoint_grad=True) == ( + 20 * 188416 + ) r._moe_gradient_shared_bytes = 188417 with pytest.raises(ValueError, match="shared-expert"): r._moe_workspace_bytes(10, checkpoint_grad=True) diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py new file mode 100644 index 000000000..a7ad7e6f1 --- /dev/null +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -0,0 +1,161 @@ +"""CP2 layout-aware recompute pricing: per-rank ledgers and gates, CPU only.""" + +from dataclasses import replace +from types import SimpleNamespace + +import pytest +from test_trainer_rank_checkpoint_memory import rank +import torch + +from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import _TE_CUBLAS_WORKSPACE_BYTES, _GroupLayout + +H = 2048 * 2 + + +def qwen36(r): + """Qwen3.6-35B-A3B's 3:1 GDN/attention pattern and mixer geometry at CP2.""" + r._topology_key = lambda: (1, 1, 2, 1) + r._geometry = replace( + r._geometry, + num_attention_heads=16, + num_query_groups=2, + kv_channels=256, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=32, + gdn_value_head_dim=128, + ) + r._attention_output_gate = True + r._gdn_layers = 30 + for index, layer in enumerate(r.runtime.model[0].decoder.layers): + is_gdn = index % 4 != 3 + after_gdn = index > 0 and (index - 1) % 4 != 3 + layer._art_gdn_island_boundary = SimpleNamespace( + is_gdn=is_gdn, input_layout="gdn" if is_gdn and after_gdn else "attention" + ) + return r + + +def test_boundaries_follow_each_layer_inputs_layout(): + # Traced real-data CP2 ranks: 20 inputs in each layout, and boundaries of + # 8.248 and 7.612 GB. The busiest attention rank is not the busiest overall. + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + layout = _GroupLayout((52480, 44314), (48194, 48600), (0, 0)) + retained, _ = r._layout_checkpoint_floor(layers, (None,), (48397,), (layout,)) + assert retained == H * (20 * 52480 + 20 * 48194) + assert H * (20 * 44314 + 20 * 48600) < retained < 40 * 52480 * H + + +def test_largest_rank_total_not_a_sum_of_rank_maxima(): + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + widths = r._recomputed_mixer_widths(stage_buffers=False) + assert widths == {"attention": 2 * (2 * 2048 + 7 * 4096 + 3 * 512), "gdn": 94208} + # Rank 1's two full-query stages keep 3.08 GB against rank 0's 0.97 GB. + layout = _GroupLayout((52480, 44314), (48194, 48600), (970_720_000, 3_079_000_000)) + retained, workspace = r._layout_checkpoint_floor( + layers, (None,), (48397,), (layout,) + ) + + def total(rank_index): + attention = layout.attention_rows[rank_index] + gdn = layout.gdn_rows[rank_index] + ledger = H * (20 * attention + 20 * gdn) + stages = ( + attention * (widths["attention"] + 2 * H) + + layout.attention_retained[rank_index] + + r._moe_workspace_bytes( + attention, routed_rows=48397, checkpoint_grad=True + ), + gdn * (widths["gdn"] + 2 * H) + + r._moe_workspace_bytes(gdn, routed_rows=48397, checkpoint_grad=True), + ) + return ledger + max(stages) + + assert retained + workspace == max(total(0), total(1)) + _TE_CUBLAS_WORKSPACE_BYTES + # Rank 1 has fewer rows but is the largest total: its attention runs two + # full-query stages. Pricing it on its own rows is about neutral against + # the busiest-rank floor, which over-counts boundaries and GDN instead. + assert total(1) > total(0) + busiest = sum(r._checkpoint_memory_floor(((52480, True),), None, (48397,))) + assert abs(retained + workspace - busiest) < 0.01 * busiest + + +def test_routed_rows_above_a_ranks_own_rows_are_all_priced(): + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + few = _GroupLayout((100, 100), (100, 100), (0, 0)) + many = _GroupLayout((100, 100), (100, 100), (0, 0)) + low = sum(r._layout_checkpoint_floor(layers, (None,), (100,), (few,))) + high = sum(r._layout_checkpoint_floor(layers, (None,), (400,), (many,))) + assert high - low == 300 * 188416 + + +def test_no_grad_groups_keep_busiest_rank_pricing(): + r = qwen36(rank()) + layout = _GroupLayout((10, 8), (9, 9), (0, 0)) + groups = ((10, True), (12, False)) + assert r._checkpoint_memory_floor( + groups, None, (9, 12), (layout, layout) + ) == r._checkpoint_memory_floor(groups, None, (9, 12)) + + +def _plan(r, *, no_grad=False): + plan = r._plan_flat_forward( + [ + ForwardInput( + input_tokens=torch.arange(64), hidden_states=True, no_grad=no_grad + ) + ] + ) + return replace(plan, signature=replace(plan.signature, topology=(1, 1, 2, 1))) + + +@pytest.mark.parametrize( + "topology", [(1, 1, 1, 1), (1, 2, 2, 1), (1, 1, 4, 1), (1, 1, 2, 2)] +) +def test_layout_pricing_is_cp2_tp1_pp1_only(topology): + r = qwen36(rank()) + plan = _plan(r) + plan = replace(plan, signature=replace(plan.signature, topology=topology)) + assert r._plan_group_layouts(plan) is None + + +def test_layout_pricing_needs_gradient_groups_and_art_cp_attention(): + r = qwen36(rank()) + assert r._plan_group_layouts(_plan(r, no_grad=True)) is None + # The stub decoder's attention layers carry no ART CP core attention. + assert r._plan_group_layouts(_plan(r)) is None + # A GDN model without island boundaries is not modeled either. + for layer in r.runtime.model[0].decoder.layers: + del layer._art_gdn_island_boundary + assert r._plan_group_layouts(_plan(r)) is None + + +def test_plan_cost_and_admission_use_the_same_layouts(monkeypatch): + r = qwen36(rank()) + plan = _plan(r) + layout = _GroupLayout((40, 24), (32, 32), (1_000_000, 3_000_000)) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((40, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (32,)) + monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: 32) + cost = r._plan_cost(plan) + # Admission adds CP output coexistence the plan cost leaves out (as it did + # before layouts); the layouts themselves enter both the same way. + gap = r._memory_check(plan).estimated_required_bytes - cost.required + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: None) + assert gap == r._memory_check(plan).estimated_required_bytes - ( + r._plan_cost(plan).required + ) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + retained, _ = r._layout_checkpoint_floor( + r.runtime.model[0].decoder.layers, + (None,), + r._plan_group_routed_rows(plan), + (layout,), + ) + assert cost.checkpoint_retained == plan.output_bytes + retained diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index d4e390ccf..c44363d4c 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -606,7 +606,7 @@ def held(rows): monkeypatch.setattr( rank, "_checkpoint_memory_floor", - lambda rows, refs=None, routed=None: (10**7, 10**6), + lambda rows, refs=None, routed=None, layouts=None: (10**7, 10**6), ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) From 0e7255c8b990d38ebf5af92689ba869b160f6ae0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 23:30:55 +0000 Subject: [PATCH 03/19] Bound CP2 layout pricing from below and count flex's own LSE The split search's lower bound priced gated plans with the busiest-rank attention width, above what per-rank pricing can charge; price even-share layouts at the least attention state instead (one aligned local stage per rank). Count flex's saved LSE beside the normalized copy it returns. Share one gate between plan pricing and the bound, checking the topology before model state. Record group layouts in planner evidence. Tests cover the positive gate path against the executor's own plans, the runtime layouts, the softmax-offset gate, tiny and single-head stages, and the lower bound staying below plan cost. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/executor.py | 29 +++- src/art/trainer_rank/_impl.py | 95 +++++++++--- .../test_context_parallel_retained_bytes.py | 50 +++++-- tests/unit/test_trainer_rank_layout_memory.py | 137 +++++++++++++++++- 4 files changed, 276 insertions(+), 35 deletions(-) diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 78a65069f..857b5ef2a 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -1551,6 +1551,26 @@ def _merge_stage_output_grads_from_tape( return stage_out_grads, stage_lse_grads +def minimum_retained_bytes_per_row( + *, + q_heads: int, + kv_heads: int, + head_dim: int, + value_head_dim: int, + element_size: int, +) -> int: + """The least ``retained_stage_record_bytes`` keeps per own row. + + Every rank with rows runs a local stage over all of them (each row attends + to itself); an aligned one keeps only the contiguous copies of multi-head + views, flex's output and its two LSEs. + """ + copies = (q_heads > 1) * q_heads * head_dim + (kv_heads > 1) * kv_heads * ( + head_dim + value_head_dim + ) + return (copies + q_heads * value_head_dim) * element_size + 2 * q_heads * 4 + + def retained_stage_record_bytes( rank_plan: RankRuntimePlan, *, @@ -1569,9 +1589,10 @@ def retained_stage_record_bytes( differs, else a contiguous copy of the permuted ``q_flat``/``k_flat`` view; partial-range gathers and remote fetch buffers are kept as the stage's inputs. It keeps flex's output and LSE at the execution length, and - logical-length copies of them when padded. Every producing stage after the - first keeps a merge-tape clone of the accumulators. Accumulators themselves - are transient. + logical-length copies of them when padded, plus flex's own LSE beside the + normalized one the FLASH backend returns (counted on every backend). Every + producing stage after the first keeps a merge-tape clone of the + accumulators. Accumulators themselves are transient. """ own = int(rank_plan.local_valid_lengths[0]) if rank_plan.local_valid_lengths else 0 accum_size = 4 if element_size < 4 else element_size @@ -1615,7 +1636,7 @@ def retained_stage_record_bytes( total += (k_row + v_row) * k_pad elif k_full and kv_heads > 1: total += (k_row + v_row) * k_len - total += (out_row + lse_row) * q_pad + total += (out_row + 2 * lse_row) * q_pad if q_pad != q_len: total += (out_row + lse_row) * q_len tape = tape_row * (own if q_full else q_len) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 389a865c7..5e0e9c88a 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -3680,6 +3680,7 @@ def _split_chunk_lower_cost( unshared_packed_tokens = 0 head_workspace_bytes = 0 group_rows: list[tuple[int, bool]] = [] + group_physical_rows: list[int] = [] for (_slot, grad_enabled), group_indices in groups: estimated = estimate_prefix_tree_packed_tokens( (rows[index] for index in group_indices), @@ -3687,6 +3688,7 @@ def _split_chunk_lower_cost( ) assert estimated is not None # rows are CPU copies physical_rows = self._physical_tokens(estimated) + group_physical_rows.append(physical_rows) packed_tokens += physical_rows # The most loaded CP rank holds at least an even share. cp = max(1, self._topology_key()[2]) @@ -3712,12 +3714,23 @@ def _split_chunk_lower_cost( slot_groups=tuple(key for key, _ in groups), ) logical_tokens = _active_logical_tokens(requests) + # Where exact plans price each rank's layouts, bound them from below + # the same way rather than with the busiest-rank widths. + layouts = ( + self._minimum_layouts(group_physical_rows, signature.topology[2]) + if self._layout_pricing_supported( + signature.topology, + gradient_groups=all(grad for _, grad in group_rows), + ) + else None + ) cost = self._subforward_cost( packed_tokens=packed_tokens, output_bytes=output_bytes, signature=signature, logical_tokens=logical_tokens, group_rows=tuple(group_rows), + group_layouts=layouts, slot_refs=tuple(ref for (ref, _), _ in groups), head_workspace_bytes=head_workspace_bytes, # The average CP load is an optimistic bound, not an admission cost. @@ -4356,37 +4369,29 @@ def _layout_checkpoint_floor( workspace += self._te_workspace_growth_bytes() return retained, workspace - def _plan_group_layouts( - self, plan: _FlatForwardPlan - ) -> tuple[_GroupLayout, ...] | None: - """Every rank's CP layouts per group, where layout pricing is modeled. - - Only CP2 at TP1/PP1 with gradient groups, ART's CP core attention with - no softmax offset, and GDN layers marked with island boundaries; the - executor's retained set is validated there. Elsewhere ``None`` keeps - the busiest-rank pricing. - """ - _dp, tp, cp, pp = plan.signature.topology - if (tp, cp, pp) != (1, 2, 1) or not plan.groups: - return None - if not all(group.grad_enabled for group in plan.groups): - return None + def _layout_pricing_supported( + self, topology: tuple[int, int, int, int], *, gradient_groups: bool + ) -> bool: + """Whether layout-aware pricing models this runtime and plan shape.""" + _dp, tp, cp, pp = topology + if (tp, cp, pp) != (1, 2, 1) or not gradient_groups: + return False geometry = self._geometry if not geometry.num_attention_heads or not geometry.kv_channels: - return None + return False try: decoder = _language_model(self.runtime.model[0]).decoder from art.megatron.context_parallel.core_attention import ( ArtContextParallelCoreAttention, ) except (AttributeError, RuntimeError, ModuleNotFoundError): - return None + return False for layer in decoder.layers: boundary = getattr(layer, "_art_gdn_island_boundary", None) if boundary is not None and boundary.is_gdn: continue if self._gdn_layers and boundary is None: - return None + return False core = getattr( getattr(layer, "self_attention", None), "core_attention", None ) @@ -4394,7 +4399,54 @@ def _plan_group_layouts( type(core) is not ArtContextParallelCoreAttention or getattr(core, "softmax_offset", None) is not None ): - return None + return False + return True + + def _minimum_layouts( + self, physical_rows: Sequence[int], cp: int + ) -> tuple[_GroupLayout, ...]: + """Even-share layouts keeping the least attention state: a lower bound. + + Some rank holds at least an even share of each layout's rows and runs + at least one aligned local stage over them. + """ + from art.megatron.context_parallel.executor import ( + minimum_retained_bytes_per_row, + ) + + geometry = self._geometry + per_row = minimum_retained_bytes_per_row( + q_heads=int(geometry.num_attention_heads), + kv_heads=int(geometry.num_query_groups), + head_dim=int(geometry.kv_channels), + value_head_dim=int(geometry.kv_channels), + element_size=self._param_dtype_size, + ) + return tuple( + _GroupLayout( + attention_rows=(rows // cp,) * cp, + gdn_rows=(rows // cp,) * cp if self._gdn_layers else None, + attention_retained=(rows // cp * per_row,) * cp, + ) + for rows in physical_rows + ) + + def _plan_group_layouts( + self, plan: _FlatForwardPlan + ) -> tuple[_GroupLayout, ...] | None: + """Every rank's CP layouts per group, where layout pricing is modeled. + + Only CP2 at TP1/PP1 with gradient groups, ART's CP core attention with + no softmax offset, and GDN layers marked with island boundaries; the + executor's retained set is validated there. Elsewhere ``None`` keeps + the busiest-rank pricing. + """ + if not plan.groups or not self._layout_pricing_supported( + plan.signature.topology, + gradient_groups=all(group.grad_enabled for group in plan.groups), + ): + return None + geometry = self._geometry from art.megatron.context_parallel.executor import retained_stage_record_bytes from art.megatron.context_parallel.runtime import context_parallel_rank_layouts from art.megatron.flex_attn.compiled import flash_sparse_block_size_for_head_dim @@ -6675,6 +6727,11 @@ def _fill_planner_snapshot( "retained_tokens": self._plan_retained_tokens(child), "group_rows": self._plan_group_rows(child), "group_routed_rows": self._plan_group_routed_rows(child), + "group_layouts": ( + None + if (layouts := self._plan_group_layouts(child)) is None + else [asdict(layout) for layout in layouts] + ), "hybridep_growth_bytes": ( self._plan_hybridep_growth_bytes(child) ), diff --git a/tests/unit/test_context_parallel_retained_bytes.py b/tests/unit/test_context_parallel_retained_bytes.py index fd5e0e643..a12b88a0c 100644 --- a/tests/unit/test_context_parallel_retained_bytes.py +++ b/tests/unit/test_context_parallel_retained_bytes.py @@ -5,6 +5,7 @@ pytest.importorskip("triton") from art.megatron.context_parallel.executor import ( # noqa: E402 + minimum_retained_bytes_per_row, retained_stage_record_bytes, ) from art.megatron.context_parallel.types import ( # noqa: E402 @@ -23,7 +24,9 @@ element_size=2, block_size=(128, 64), ) -Q, KV, OUT, TAPE = 8192, 2048, 8192 + 64, 16 * 257 * 4 +# Flex's output keeps its own LSE beside the normalized copy it returns; +# logical-length copies keep one. +Q, KV, FLEX, OUT, TAPE = 8192, 2048, 8192 + 2 * 64, 8192 + 64, 16 * 257 * 4 def stage(index, *, local, q, k, q_len=None, k_len=None, own=None, source=0): @@ -61,16 +64,41 @@ def test_aligned_single_stage_copies_views_without_output_copies(): retained = retained_stage_record_bytes( plan(rows, stage(0, local=True, q=rows, k=rows)), **GEOMETRY ) - assert retained == Q * rows + KV * rows + OUT * rows # 0.971 GB traced + assert retained == Q * rows + KV * rows + FLEX * rows # 0.971 GB traced + geometry = {k: v for k, v in GEOMETRY.items() if k != "block_size"} + assert retained == rows * minimum_retained_bytes_per_row(**geometry) def test_unaligned_single_stage_pads_and_copies_the_logical_output(): + # The tail chunk on the single-stage rank: the planner rounds both lengths + # up to its 128-row block, so both pad. rows = 52481 - q_pad, k_pad = 411 * 128, 821 * 64 + stage_len = 411 * 128 retained = retained_stage_record_bytes( - plan(rows, stage(0, local=True, q=rows, k=rows)), **GEOMETRY + plan( + rows, stage(0, local=True, q=rows, k=rows, q_len=stage_len, k_len=stage_len) + ), + **GEOMETRY, + ) + assert retained == Q * stage_len + KV * stage_len + FLEX * stage_len + OUT * rows + + +def test_tiny_stage_pads_to_two_blocks(): + retained = retained_stage_record_bytes( + plan(5, stage(0, local=True, q=5, k=5)), **GEOMETRY + ) + assert retained == Q * 256 + KV * 128 + FLEX * 256 + OUT * 5 + + +def test_single_head_views_are_already_contiguous(): + rows = 1024 + geometry = dict(GEOMETRY, q_heads=1, kv_heads=1) + retained = retained_stage_record_bytes( + plan(rows, stage(0, local=True, q=rows, k=rows)), **geometry ) - assert retained == Q * q_pad + KV * k_pad + OUT * q_pad + OUT * rows + assert retained == (256 * 2 + 2 * 4) * rows + del geometry["block_size"] + assert minimum_retained_bytes_per_row(**geometry) == 256 * 2 + 2 * 4 def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): @@ -81,9 +109,9 @@ def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): remote = stage(1, local=False, q=own, k=remote_k, q_len=44352, k_len=16512) retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) q_pad = 347 * 128 # 44,416, as flex's traced output size shows - local_bytes = Q * q_pad + KV * 44352 + OUT * q_pad + OUT * own + local_bytes = Q * q_pad + KV * 44352 + FLEX * q_pad + OUT * own remote_bytes = ( - Q * q_pad + KV * remote_k + KV * 16512 + OUT * q_pad + OUT * own + TAPE * own + Q * q_pad + KV * remote_k + KV * 16512 + FLEX * q_pad + OUT * own + TAPE * own ) assert retained == local_bytes + remote_bytes assert 3.05e9 < retained < 3.10e9 # 3.07 GB traced @@ -95,9 +123,9 @@ def test_partial_query_remote_stage_keeps_its_gather_and_a_partial_tape(): local = stage(0, local=True, q=own, k=own, q_len=105216, k_len=105216) remote = stage(1, local=False, q=768, k=20608, own=own) retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) - local_bytes = Q * 105216 + KV * 105216 + OUT * 105216 + OUT * own + local_bytes = Q * 105216 + KV * 105216 + FLEX * 105216 + OUT * own # Aligned: the query gather and fetch buffers feed flex without copies. - remote_bytes = Q * 768 + KV * 20608 + OUT * 768 + TAPE * 768 + remote_bytes = Q * 768 + KV * 20608 + FLEX * 768 + TAPE * 768 assert retained == local_bytes + remote_bytes @@ -107,12 +135,12 @@ def test_empty_remote_stage_and_missing_local_stage(): alone = retained_stage_record_bytes( plan(rows, stage(0, local=True, q=rows, k=rows), empty), **GEOMETRY ) - assert alone == Q * rows + KV * rows + OUT * rows + assert alone == Q * rows + KV * rows + FLEX * rows # Without a local stage, the first ready remote stage records no tape; the # mirror drops the smallest so it never under-counts the order. small = stage(1, local=False, q=256, k=256, own=rows) full = stage(2, local=False, q=rows, k=512) both = retained_stage_record_bytes(plan(rows, small, full), **GEOMETRY) without_small_tape = retained_stage_record_bytes(plan(rows, full), **GEOMETRY) - assert both - without_small_tape == Q * 256 + KV * 256 + OUT * 256 + TAPE * rows + assert both - without_small_tape == Q * 256 + KV * 256 + FLEX * 256 + TAPE * rows assert retained_stage_record_bytes(plan(rows), **GEOMETRY) == 0 diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index a7ad7e6f1..37030533d 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -8,7 +8,7 @@ import torch from art.trainer_rank import ForwardInput -from art.trainer_rank._impl import _TE_CUBLAS_WORKSPACE_BYTES, _GroupLayout +from art.trainer_rank._impl import _TE_CUBLAS_WORKSPACE_BYTES, Unset, _GroupLayout H = 2048 * 2 @@ -159,3 +159,138 @@ def test_plan_cost_and_admission_use_the_same_layouts(monkeypatch): (layout,), ) assert cost.checkpoint_retained == plan.output_bytes + retained + + +def art_cp(r, monkeypatch): + """Qwen3.6 at CP2 with ART's CP core attention and a CPU planning config.""" + from art.megatron.context_parallel.core_attention import ( + ArtContextParallelCoreAttention, + ) + from art.megatron.context_parallel.types import ParallelTopology + + r = qwen36(r) + for index, layer in enumerate(r.runtime.model[0].decoder.layers): + if index % 4 == 3: + core = ArtContextParallelCoreAttention.__new__( + ArtContextParallelCoreAttention + ) + torch.nn.Module.__init__(core) + core.softmax_offset = None + layer.self_attention = torch.nn.Module() + layer.self_attention.core_attention = core + provider = r.runtime.provider + provider.kv_channels = 256 + provider.ffn_hidden_size = 8192 + provider.linear_num_key_heads = 16 + provider.linear_num_value_heads = 32 + provider.linear_key_head_dim = 128 + provider.linear_value_head_dim = 128 + provider.params_dtype = torch.bfloat16 + r.runtime.model_support_handler = SimpleNamespace( + build_gdn_execution_spec=True, + context_parallel_workload_profile=lambda provider: None, + ) + monkeypatch.setattr(r, "_topology", lambda: ParallelTopology(tp=1, cp=2)) + return r + + +def _requests(lengths=(900, 700, 500)): + start, requests = 0, [] + for n in lengths: + requests.append( + ForwardInput( + input_tokens=torch.arange(start, start + n), hidden_states=True + ) + ) + start += n + return requests + + +def test_rank_layouts_are_the_executors_plans(monkeypatch): + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.context_parallel.types import ( + ContextParallelConfig, + ParallelTopology, + ) + from art.megatron.prefix_tree_packing import prefix_tree_pack + + packed = prefix_tree_pack([r.input_tokens for r in _requests()], max_depth=1) + attention, gdn, plans = context_parallel_rank_layouts( + group_ids=packed.group_ids, + parent_ids=packed.parent_ids, + topology=ParallelTopology(tp=1, cp=2), + config=ContextParallelConfig(), + original_seq_len=int(packed.tokens.shape[1]), + build_gdn_execution_spec=True, + ) + total = int(packed.tokens.numel()) + assert sum(attention) == total and gdn is not None and sum(gdn) == total + # The ledger's attention rows and the mirror's own rows are the same count. + assert attention == tuple(plan.local_valid_lengths[0] for plan in plans) + + +def test_gated_plans_price_every_rank_from_its_stage_plan(monkeypatch): + from art.megatron.context_parallel.executor import retained_stage_record_bytes + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.context_parallel.types import ParallelTopology + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + ) + + r = art_cp(rank(), monkeypatch) + plan = _plan_with(r, _requests()) + (layout,) = r._plan_group_layouts(plan) + assert sum(layout.attention_rows) == plan.packed_tokens + assert layout.gdn_rows is not None and sum(layout.gdn_rows) == plan.packed_tokens + # Each rank's retention is the executor mirror over that rank's own plan. + (group,) = plan.groups + _, _, rank_plans = context_parallel_rank_layouts( + group_ids=group.packed.group_ids, + parent_ids=group.packed.parent_ids, + topology=ParallelTopology(tp=1, cp=2), + config=_context_parallel_config_for_provider( + r.runtime.provider, r.device, r.runtime.model_support_handler + ), + original_seq_len=int(group.packed.tokens.shape[1]), + build_gdn_execution_spec=True, + ) + assert layout.attention_retained == tuple( + retained_stage_record_bytes( + rank_plan, + q_heads=16, + kv_heads=2, + head_dim=256, + value_head_dim=256, + element_size=2, + block_size=(128, 128), # the CPU device's flex block + ) + for rank_plan in rank_plans + ) + assert all(retained > 0 for retained in layout.attention_retained) + + +def test_softmax_offset_leaves_the_busiest_rank_floor(monkeypatch): + r = art_cp(rank(), monkeypatch) + plan = _plan_with(r, _requests()) + assert r._plan_group_layouts(plan) is not None + for layer in r.runtime.model[0].decoder.layers: + core = getattr(getattr(layer, "self_attention", None), "core_attention", None) + if core is not None: + core.softmax_offset = torch.zeros(16) + assert r._plan_group_layouts(plan) is None + + +def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch): + r = art_cp(rank(), monkeypatch) + requests = _requests((2048, 1536, 1024, 512)) + plan = _plan_with(r, requests) + lower = r._split_chunk_lower_cost( + requests, tuple(item.input_tokens for item in requests), checkpoint=Unset + ) + assert lower.required <= r._plan_cost(plan).required + # It prices even-share layouts at the least attention state. + assert r._layout_pricing_supported((1, 1, 2, 1), gradient_groups=True) + + +def _plan_with(r, requests): + return r._plan_flat_forward(requests) From 9b4e9d709691b0dfdc33087761d34fee9dcf2307 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 23:43:04 +0000 Subject: [PATCH 04/19] Guard the layout gate and state why the lower bound rounds down The gate now declines models with several chunks or a decoder without layers before reading them. The bound holds because every rank's total grows with its own rows, so the largest is at least the total at the mean; rounding the even share up can exceed a split's exact cost. The lower-bound test covers even, skewed and odd splits and one long sequence. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 14 +++++++++++--- tests/unit/test_trainer_rank_layout_memory.py | 9 +++++++-- 2 files changed, 18 insertions(+), 5 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5e0e9c88a..bc47b03ed 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4379,6 +4379,8 @@ def _layout_pricing_supported( geometry = self._geometry if not geometry.num_attention_heads or not geometry.kv_channels: return False + if len(self.runtime.model) != 1: + return False try: decoder = _language_model(self.runtime.model[0]).decoder from art.megatron.context_parallel.core_attention import ( @@ -4386,7 +4388,10 @@ def _layout_pricing_supported( ) except (AttributeError, RuntimeError, ModuleNotFoundError): return False - for layer in decoder.layers: + layers = getattr(decoder, "layers", None) + if layers is None: + return False + for layer in layers: boundary = getattr(layer, "_art_gdn_island_boundary", None) if boundary is not None and boundary.is_gdn: continue @@ -4407,8 +4412,11 @@ def _minimum_layouts( ) -> tuple[_GroupLayout, ...]: """Even-share layouts keeping the least attention state: a lower bound. - Some rank holds at least an even share of each layout's rows and runs - at least one aligned local stage over them. + Every rank's total grows with its own rows, and every rank keeps at + least an aligned local stage's state per row (each row attends to + itself), so the largest rank's total is at least the total at the mean + rows. The mean is at least the floor of an even share, which is why + this rounds down; rounding up can exceed a split's exact cost. """ from art.megatron.context_parallel.executor import ( minimum_retained_bytes_per_row, diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index 37030533d..baf88257b 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -280,9 +280,14 @@ def test_softmax_offset_leaves_the_busiest_rank_floor(monkeypatch): assert r._plan_group_layouts(plan) is None -def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch): +@pytest.mark.parametrize( + "lengths", + [(2048, 1536, 1024, 512), (4099, 3, 5, 7), (1, 2, 3, 4, 5, 6, 7), (8191,)], +) +def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths): + # Even and skewed CP splits, odd row counts, and a single long sequence. r = art_cp(rank(), monkeypatch) - requests = _requests((2048, 1536, 1024, 512)) + requests = _requests(lengths) plan = _plan_with(r, requests) lower = r._split_chunk_lower_cost( requests, tuple(item.input_tokens for item in requests), checkpoint=Unset From 33185b49ed11ccbff0f5c433d5b8cee792e4e3d3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 23:52:43 +0000 Subject: [PATCH 05/19] Run the CP layout memory tests in the Megatron CI lane Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 4ff185190..b66ac2c0e 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -239,6 +239,8 @@ jobs: tests/unit/test_trainer_rank_pending_memory.py \ tests/unit/test_trainer_rank_shared_memory.py \ tests/unit/test_trainer_rank_converted_memory.py \ + tests/unit/test_trainer_rank_layout_memory.py \ + tests/unit/test_context_parallel_retained_bytes.py \ tests/unit/test_trainer_rank_split.py \ tests/unit/test_megatron_compile_garbage.py \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ @@ -283,4 +285,6 @@ jobs: --ignore=tests/unit/test_trainer_rank_pending_memory.py \ --ignore=tests/unit/test_trainer_rank_shared_memory.py \ --ignore=tests/unit/test_megatron_compile_garbage.py \ - --ignore=tests/unit/test_trainer_rank_converted_memory.py + --ignore=tests/unit/test_trainer_rank_converted_memory.py \ + --ignore=tests/unit/test_trainer_rank_layout_memory.py \ + --ignore=tests/unit/test_context_parallel_retained_bytes.py From c1500cafbb98cde712201f97f9768b63f9511445 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 00:24:13 +0000 Subject: [PATCH 06/19] Charge single-head view copies and partial-tape indices; time layout planning One head does not make a view of a fused QKV split contiguous, so the mirror now charges the copy flex makes of any full local view; the lower bound still counts only multi-head copies. A partial-query merge tape also keeps its int64 row index. The per-rank layout work now counts toward planning time (about 40-76 ms for a new layout's peer plan, under 1 ms once cached). Tests use typed geometry helpers. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/executor.py | 10 +- src/art/trainer_rank/_impl.py | 10 ++ .../test_context_parallel_retained_bytes.py | 92 ++++++++++--------- tests/unit/test_trainer_rank_layout_memory.py | 1 + 4 files changed, 68 insertions(+), 45 deletions(-) diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 857b5ef2a..5d1295963 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -1623,7 +1623,10 @@ def retained_stage_record_bytes( total += q_row * q_len if q_pad != q_len: total += q_row * q_pad - elif q_full and q_heads > 1: + elif q_full: + # A view of the projection output: flex copies it unless it is + # already contiguous, which one head alone does not guarantee + # (a fused QKV split keeps the projection's token stride). total += q_row * q_len # Keys and values: local ranges as for queries; remote ones land in # contiguous head-major fetch buffers kept as the stage's inputs. @@ -1634,12 +1637,13 @@ def retained_stage_record_bytes( total += (k_row + v_row) * k_len if k_pad != k_len: total += (k_row + v_row) * k_pad - elif k_full and kv_heads > 1: + elif k_full: total += (k_row + v_row) * k_len total += (out_row + 2 * lse_row) * q_pad if q_pad != q_len: total += (out_row + lse_row) * q_len - tape = tape_row * (own if q_full else q_len) + # A partial-query tape also keeps its int64 row index. + tape = tape_row * own if q_full else (tape_row + 8) * q_len if stage.is_local_stage: local_produced = True else: diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index bc47b03ed..6d3d61b98 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4454,6 +4454,16 @@ def _plan_group_layouts( gradient_groups=all(group.grad_enabled for group in plan.groups), ): return None + started = time.perf_counter() + try: + return self._compute_group_layouts(plan) + finally: + # Planning work: every rank's CP plan, cached by planning key. + self._planning_seconds_accum += time.perf_counter() - started + + def _compute_group_layouts( + self, plan: _FlatForwardPlan + ) -> tuple[_GroupLayout, ...]: geometry = self._geometry from art.megatron.context_parallel.executor import retained_stage_record_bytes from art.megatron.context_parallel.runtime import context_parallel_rank_layouts diff --git a/tests/unit/test_context_parallel_retained_bytes.py b/tests/unit/test_context_parallel_retained_bytes.py index a12b88a0c..583de3a24 100644 --- a/tests/unit/test_context_parallel_retained_bytes.py +++ b/tests/unit/test_context_parallel_retained_bytes.py @@ -14,16 +14,31 @@ TokenRange, ) + # Qwen3.6-35B-A3B attention: 16 query heads, 2 KV heads of 256, BF16, and the # H200 flash block for a 256-wide head (128 query, 64 key rows). -GEOMETRY = dict( - q_heads=16, - kv_heads=2, - head_dim=256, - value_head_dim=256, - element_size=2, - block_size=(128, 64), -) +def retained(runtime_plan, *, q_heads=16, kv_heads=2): + return retained_stage_record_bytes( + runtime_plan, + q_heads=q_heads, + kv_heads=kv_heads, + head_dim=256, + value_head_dim=256, + element_size=2, + block_size=(128, 64), + ) + + +def minimum(*, q_heads=16, kv_heads=2): + return minimum_retained_bytes_per_row( + q_heads=q_heads, + kv_heads=kv_heads, + head_dim=256, + value_head_dim=256, + element_size=2, + ) + + # Flex's output keeps its own LSE beside the normalized copy it returns; # logical-length copies keep one. Q, KV, FLEX, OUT, TAPE = 8192, 2048, 8192 + 2 * 64, 8192 + 64, 16 * 257 * 4 @@ -61,12 +76,9 @@ def test_aligned_single_stage_copies_views_without_output_copies(): # Real-data rank 0: one aligned local stage. Contiguous copies of the # permuted Q/K/V views, flex output and LSE; no padding, copies or tape. rows = 52480 - retained = retained_stage_record_bytes( - plan(rows, stage(0, local=True, q=rows, k=rows)), **GEOMETRY - ) - assert retained == Q * rows + KV * rows + FLEX * rows # 0.971 GB traced - geometry = {k: v for k, v in GEOMETRY.items() if k != "block_size"} - assert retained == rows * minimum_retained_bytes_per_row(**geometry) + kept = retained(plan(rows, stage(0, local=True, q=rows, k=rows))) + assert kept == Q * rows + KV * rows + FLEX * rows # 0.971 GB traced + assert kept == rows * minimum() def test_unaligned_single_stage_pads_and_copies_the_logical_output(): @@ -74,31 +86,28 @@ def test_unaligned_single_stage_pads_and_copies_the_logical_output(): # up to its 128-row block, so both pad. rows = 52481 stage_len = 411 * 128 - retained = retained_stage_record_bytes( + kept = retained( plan( rows, stage(0, local=True, q=rows, k=rows, q_len=stage_len, k_len=stage_len) - ), - **GEOMETRY, + ) ) - assert retained == Q * stage_len + KV * stage_len + FLEX * stage_len + OUT * rows + assert kept == Q * stage_len + KV * stage_len + FLEX * stage_len + OUT * rows def test_tiny_stage_pads_to_two_blocks(): - retained = retained_stage_record_bytes( - plan(5, stage(0, local=True, q=5, k=5)), **GEOMETRY - ) - assert retained == Q * 256 + KV * 128 + FLEX * 256 + OUT * 5 + kept = retained(plan(5, stage(0, local=True, q=5, k=5))) + assert kept == Q * 256 + KV * 128 + FLEX * 256 + OUT * 5 -def test_single_head_views_are_already_contiguous(): +def test_single_head_views_may_still_be_copied(): + # One head does not make a view of a fused QKV split contiguous, so the + # mirror still charges the copies; only the lower bound leaves them out. rows = 1024 - geometry = dict(GEOMETRY, q_heads=1, kv_heads=1) - retained = retained_stage_record_bytes( - plan(rows, stage(0, local=True, q=rows, k=rows)), **geometry + kept = retained( + plan(rows, stage(0, local=True, q=rows, k=rows)), q_heads=1, kv_heads=1 ) - assert retained == (256 * 2 + 2 * 4) * rows - del geometry["block_size"] - assert minimum_retained_bytes_per_row(**geometry) == 256 * 2 + 2 * 4 + assert kept == (512 + 1024 + 512 + 2 * 4) * rows + assert minimum(q_heads=1, kv_heads=1) == 512 + 2 * 4 def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): @@ -107,14 +116,14 @@ def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): own, remote_k = 44314, 16504 local = stage(0, local=True, q=own, k=own, q_len=44352, k_len=44352) remote = stage(1, local=False, q=own, k=remote_k, q_len=44352, k_len=16512) - retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) + kept = retained(plan(own, local, remote)) q_pad = 347 * 128 # 44,416, as flex's traced output size shows local_bytes = Q * q_pad + KV * 44352 + FLEX * q_pad + OUT * own remote_bytes = ( Q * q_pad + KV * remote_k + KV * 16512 + FLEX * q_pad + OUT * own + TAPE * own ) - assert retained == local_bytes + remote_bytes - assert 3.05e9 < retained < 3.10e9 # 3.07 GB traced + assert kept == local_bytes + remote_bytes + assert 3.05e9 < kept < 3.10e9 # 3.07 GB traced def test_partial_query_remote_stage_keeps_its_gather_and_a_partial_tape(): @@ -122,25 +131,24 @@ def test_partial_query_remote_stage_keeps_its_gather_and_a_partial_tape(): own = 105153 local = stage(0, local=True, q=own, k=own, q_len=105216, k_len=105216) remote = stage(1, local=False, q=768, k=20608, own=own) - retained = retained_stage_record_bytes(plan(own, local, remote), **GEOMETRY) + kept = retained(plan(own, local, remote)) local_bytes = Q * 105216 + KV * 105216 + FLEX * 105216 + OUT * own - # Aligned: the query gather and fetch buffers feed flex without copies. - remote_bytes = Q * 768 + KV * 20608 + FLEX * 768 + TAPE * 768 - assert retained == local_bytes + remote_bytes + # Aligned: the query gather and fetch buffers feed flex without copies; the + # partial tape also keeps its int64 row index. + remote_bytes = Q * 768 + KV * 20608 + FLEX * 768 + (TAPE + 8) * 768 + assert kept == local_bytes + remote_bytes def test_empty_remote_stage_and_missing_local_stage(): rows = 1024 empty = stage(1, local=False, q=0, k=0) - alone = retained_stage_record_bytes( - plan(rows, stage(0, local=True, q=rows, k=rows), empty), **GEOMETRY - ) + alone = retained(plan(rows, stage(0, local=True, q=rows, k=rows), empty)) assert alone == Q * rows + KV * rows + FLEX * rows # Without a local stage, the first ready remote stage records no tape; the # mirror drops the smallest so it never under-counts the order. small = stage(1, local=False, q=256, k=256, own=rows) full = stage(2, local=False, q=rows, k=512) - both = retained_stage_record_bytes(plan(rows, small, full), **GEOMETRY) - without_small_tape = retained_stage_record_bytes(plan(rows, full), **GEOMETRY) + both = retained(plan(rows, small, full)) + without_small_tape = retained(plan(rows, full)) assert both - without_small_tape == Q * 256 + KV * 256 + FLEX * 256 + TAPE * rows - assert retained_stage_record_bytes(plan(rows), **GEOMETRY) == 0 + assert retained(plan(rows)) == 0 diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index baf88257b..0cc147c6d 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -61,6 +61,7 @@ def test_largest_rank_total_not_a_sum_of_rank_maxima(): def total(rank_index): attention = layout.attention_rows[rank_index] + assert layout.gdn_rows is not None gdn = layout.gdn_rows[rank_index] ledger = H * (20 * attention + 20 * gdn) stages = ( From 4368998e490111574164a1ceb0b8d05b638a072f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 00:30:13 +0000 Subject: [PATCH 07/19] Keep partial-tape indices when the first stage drops its tape The executor keeps a partial stage's int64 row index even when that stage produced first and recorded no accumulator tape, so charge indices apart from the tape the mirror drops. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/megatron/context_parallel/executor.py | 10 ++++++---- tests/unit/test_context_parallel_retained_bytes.py | 5 ++++- 2 files changed, 10 insertions(+), 5 deletions(-) diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 5d1295963..c59751966 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -1562,8 +1562,8 @@ def minimum_retained_bytes_per_row( """The least ``retained_stage_record_bytes`` keeps per own row. Every rank with rows runs a local stage over all of them (each row attends - to itself); an aligned one keeps only the contiguous copies of multi-head - views, flex's output and its two LSEs. + to itself); an aligned one keeps at least the contiguous copies of + multi-head views, flex's output and its two LSEs. """ copies = (q_heads > 1) * q_heads * head_dim + (kv_heads > 1) * kv_heads * ( head_dim + value_head_dim @@ -1642,8 +1642,10 @@ def retained_stage_record_bytes( total += (out_row + 2 * lse_row) * q_pad if q_pad != q_len: total += (out_row + lse_row) * q_len - # A partial-query tape also keeps its int64 row index. - tape = tape_row * own if q_full else (tape_row + 8) * q_len + tape = tape_row * (own if q_full else q_len) + if not q_full: + # Its int64 row index is kept even by the first producing stage. + total += 8 * q_len if stage.is_local_stage: local_produced = True else: diff --git a/tests/unit/test_context_parallel_retained_bytes.py b/tests/unit/test_context_parallel_retained_bytes.py index 583de3a24..18f0f4b48 100644 --- a/tests/unit/test_context_parallel_retained_bytes.py +++ b/tests/unit/test_context_parallel_retained_bytes.py @@ -150,5 +150,8 @@ def test_empty_remote_stage_and_missing_local_stage(): full = stage(2, local=False, q=rows, k=512) both = retained(plan(rows, small, full)) without_small_tape = retained(plan(rows, full)) - assert both - without_small_tape == Q * 256 + KV * 256 + FLEX * 256 + TAPE * rows + # The dropped tape's int64 index stays: the executor keeps it regardless. + assert both - without_small_tape == ( + Q * 256 + KV * 256 + FLEX * 256 + 8 * 256 + TAPE * rows + ) assert retained(plan(rows)) == 0 From b4a619ec26b753a2a72a2f274607fd9a362d0c7c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:35:12 +0000 Subject: [PATCH 08/19] Test the combine-extent floor through per-rank layouts Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_moe_memory.py | 33 ++++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3644f0a90..a90b222c0 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -867,6 +867,39 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( assert not torch.cuda.is_initialized() +def test_hybridep_combine_extent_floors_the_layout_path( + hybrid_checkpoint_rank, monkeypatch +): + rank = hybrid_checkpoint_rank + groups = ((2, True),) + calls = [] + + def layout_floor(layers, refs, routed, layouts): + calls.append(layouts) + return 7, 11 + + def generic_floor(*args): + raise AssertionError("per-rank layouts must take the layout path") + + monkeypatch.setattr(rank, "_layout_checkpoint_floor", layout_floor) + monkeypatch.setattr(rank, "_generic_checkpoint_floor", generic_floor) + layouts = (object(),) + te = rank._te_workspace_growth_bytes() + # Two rows round up to four; the layout floor's small workspace loses. + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + 4 * 2048 * 2 + te, + ) + marker = torch.empty(0) + rank._pending_hybridep_graphs.append(weakref.ref(marker)) + rank._hybridep_rows_high_water = 218751 + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + 218752 * 2048 * 2 + te, + ) + assert calls == [layouts, layouts] + + @pytest.mark.parametrize("reference", ["absent", "expired", "smaller"]) def test_hybridep_high_water_needs_a_live_larger_graph( hybrid_checkpoint_rank, reference From 0f59d7c1c6d00b672f3df41a5c71a4c7ccfe47f4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sun, 27 Sep 2026 00:22:07 +0000 Subject: [PATCH 09/19] Pin that the layout path's combine floor adds no second TE growth Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_moe_memory.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index a90b222c0..1cc118bc3 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -897,7 +897,16 @@ def generic_floor(*args): 7, 218752 * 2048 * 2 + te, ) - assert calls == [layouts, layouts] + # A larger stage already carries its TE growth; the combine floor adds none. + stage = 218752 * 2048 * 2 + te + 1 + + def larger_layout_floor(layers, refs, routed, layouts): + calls.append(layouts) + return 7, stage + + monkeypatch.setattr(rank, "_layout_checkpoint_floor", larger_layout_floor) + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == (7, stage) + assert calls == [layouts, layouts, layouts] @pytest.mark.parametrize("reference", ["absent", "expired", "smaller"]) From be1db0c2cf68c47ba5967eeb04416bd7588d699a Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 06:53:20 +0000 Subject: [PATCH 10/19] Price adapter gradients against each rank's layout boundaries Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 60 +++++++++++++++++-- tests/unit/test_trainer_rank_layout_memory.py | 54 +++++++++++++++++ 2 files changed, 109 insertions(+), 5 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 5088a1507..959551440 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4923,6 +4923,48 @@ def _checkpoint_layer_boundaries(self, retained: int) -> tuple[int, ...]: """Each decoder layer's saved-boundary bytes within the floor's retained term.""" return (retained // self._num_layers,) * self._num_layers + def _checkpoint_layer_boundary_sets( + self, + retained: int, + group_rows: tuple[tuple[int, bool], ...], + layouts: tuple[_GroupLayout, ...] | None, + ) -> tuple[tuple[int, ...], ...]: + """Each rank's saved-boundary bytes per decoder layer, as the floor prices them. + + On per-rank CP layouts (``_layout_checkpoint_floor``) a layer saves its + input in the GDN layout when it follows a GDN layer in its island and in + the attention layout otherwise, on that rank's rows. Elsewhere every + layer saves an equal share of the floor's ``retained``. + """ + if ( + layouts is None + or self._topology_key()[1] > 1 + or not all(grad for _, grad in group_rows) + ): + return (self._checkpoint_layer_boundaries(retained),) + decoder = _language_model(self.runtime.model[0]).decoder + gdn_inputs = [ + getattr( + getattr(layer, "_art_gdn_island_boundary", None), "input_layout", "" + ) + == "gdn" + for layer in decoder.layers + ] + hidden = self._hidden_size * 2 + sets = [] + for rank in range(len(layouts[0].attention_rows)): + attention = sum(max(1, layout.attention_rows[rank]) for layout in layouts) + gdn = sum( + max(1, layout.attention_rows[rank]) + if layout.gdn_rows is None + else max(1, layout.gdn_rows[rank]) + for layout in layouts + ) + sets.append( + tuple(hidden * (gdn if is_gdn else attention) for is_gdn in gdn_inputs) + ) + return tuple(sets) + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4994,9 +5036,13 @@ def _subforward_cost( # forward retention, including the cold fallback above. gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) gradient_slots = self._gradient_slots(group_rows, slot_refs) + # The busiest rank's extra, over each rank's own boundaries. adapter_gradient = ( - self._checkpoint_adapter_gradient_bytes( - gradient_slots, self._checkpoint_layer_boundaries(checkpoint_retained) + max( + self._checkpoint_adapter_gradient_bytes(gradient_slots, boundaries) + for boundaries in self._checkpoint_layer_boundary_sets( + checkpoint_retained, group_rows, group_layouts + ) ) if gradient else 0 @@ -8834,9 +8880,13 @@ def _estimate_required_memory_bytes_from_values( # as _subforward_cost. backward = self._checkpoint_input_gradient_bytes( group_rows, slot_refs - ) + self._checkpoint_adapter_gradient_bytes( - self._gradient_slots(group_rows, slot_refs), - self._checkpoint_layer_boundaries(retained), + ) + max( + self._checkpoint_adapter_gradient_bytes( + self._gradient_slots(group_rows, slot_refs), boundaries + ) + for boundaries in self._checkpoint_layer_boundary_sets( + retained, group_rows, group_layouts + ) ) if profiled is None and any(grad for _, grad in group_rows): backward += _COLD_RECOMPUTE_TRANSIENT_BYTES diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index 12942e7a3..38b52319e 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -302,3 +302,57 @@ def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths): def _plan_with(r, requests): return r._plan_flat_forward(requests) + + +def test_adapter_gradients_meet_each_ranks_own_layer_boundaries(monkeypatch): + from art.megatron.lora import LoRASlotRef + + r = qwen36(rank()) + plan = _plan(r) + layout = _GroupLayout((400, 240), (160, 320), (0, 0)) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((400, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (320,)) + monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: 320) + policy = LoRASlotRef("checkpoint", "policy") + monkeypatch.setattr( + r, "_gradient_slots", lambda group_rows, slot_refs: frozenset({policy}) + ) + pending = (300 * H,) * 40 + (0,) + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: pending if tuple(refs) else (), + ) + layers = r.runtime.model[0].decoder.layers + gdn_inputs = [ + layer._art_gdn_island_boundary.input_layout == "gdn" for layer in layers + ] + + def boundaries(rank_index): + assert layout.gdn_rows is not None + return [ + H * (layout.gdn_rows if is_gdn else layout.attention_rows)[rank_index] + for is_gdn in gdn_inputs + ] + + def extra(saved): + return max( + 0, + *( + sum(pending[i:40]) + pending[40] - sum(saved[i + 1 :]) + for i in range(40) + ), + ) + + # Each rank releases its own boundaries in layer order; the busiest extra wins. + expected = max(extra(boundaries(0)), extra(boundaries(1))) + assert r._plan_cost(plan).checkpoint_adapter_gradient == expected > 0 + # An even share of the busiest rank's boundaries would misplace the peak. + retained, _ = r._layout_checkpoint_floor(layers, (None,), (320,), (layout,)) + assert extra([retained // 40] * 40) != expected + assert r._checkpoint_layer_boundary_sets(retained, ((400, True),), (layout,)) == ( + tuple(boundaries(0)), + tuple(boundaries(1)), + ) From 42a2e2bebeb4759dbdc53d555c229a99cabb3faa Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 06:56:59 +0000 Subject: [PATCH 11/19] Pin TE growth and a stage between the combine floor and its TE Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_moe_memory.py | 25 +++++++++++++++------- 1 file changed, 17 insertions(+), 8 deletions(-) diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index c937b149c..0142e5eda 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -904,16 +904,25 @@ def generic_floor(*args): 7, 218752 * 2048 * 2 + te, ) - # A larger stage already carries its TE growth; the combine floor adds none. - stage = 218752 * 2048 * 2 + te + 1 + # The stage already carries its TE growth; the combine floor adds its own + # only when it wins, including over a stage between the two. + assert te > 0 + combine = 218752 * 2048 * 2 + for stage, workspace in ( + (combine + te - 1, combine + te), + (combine + te + 1, combine + te + 1), + ): - def larger_layout_floor(layers, refs, routed, layouts): - calls.append(layouts) - return 7, stage + def stage_layout_floor(layers, refs, routed, layouts, stage=stage): + calls.append(layouts) + return 7, stage - monkeypatch.setattr(rank, "_layout_checkpoint_floor", larger_layout_floor) - assert rank._checkpoint_memory_floor(groups, layouts=layouts) == (7, stage) - assert calls == [layouts, layouts, layouts] + monkeypatch.setattr(rank, "_layout_checkpoint_floor", stage_layout_floor) + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + workspace, + ) + assert calls == [layouts] * 4 @pytest.mark.parametrize("reference", ["absent", "expired", "smaller"]) From 0606917e834bd1705e15b42648b4ab366960299e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 07:16:21 +0000 Subject: [PATCH 12/19] Pair each CP rank's adapter-gradient extra with its own floor Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 134 +++++++++++++----- tests/unit/test_trainer_rank_layout_memory.py | 72 ++++++++-- 2 files changed, 152 insertions(+), 54 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 959551440..cef60580b 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4501,6 +4501,21 @@ def _layout_checkpoint_floor( routed: tuple[int | None, ...], layouts: tuple[_GroupLayout, ...], ) -> tuple[int, int]: + """The largest rank's boundaries and the rest of the largest rank total. + + The two sum to the largest of ``_layout_checkpoint_rank_floors``' totals. + """ + floors = self._layout_checkpoint_rank_floors(layers, refs, routed, layouts) + retained = max(retained for retained, _ in floors) + return retained, max(map(sum, floors)) - retained + + def _layout_checkpoint_rank_floors( + self, + layers: Sequence[torch.nn.Module], + refs: tuple["LoRASlotRef | None", ...], + routed: tuple[int | None, ...], + layouts: tuple[_GroupLayout, ...], + ) -> tuple[tuple[int, int], ...]: """Each CP rank's boundaries and recomputed layer on its own layouts. A saved layer input arrives in the GDN layout when the layer follows a @@ -4510,8 +4525,7 @@ def _layout_checkpoint_floor( backward; a GDN layer keeps the GDN width on its GDN rows. Either one's residual, norm, routing state and MoE stage use that layer's rows, and routed rows the EP share. Traced CP2 ranks: boundaries match exactly and - attention within 0.3%. Returns the largest rank's boundaries and the - rest of the largest rank total, so the two sum to that total. + attention within 0.3%. Returns each rank's (boundaries, workspace). """ hidden = self._hidden_size * 2 gdn_inputs = sum( @@ -4525,8 +4539,7 @@ def _layout_checkpoint_floor( widths = self._recomputed_mixer_widths(stage_buffers=False) moe = self._checkpoint_moe_bytes_per_token() beside = 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe else 0 - retained_by_rank: list[int] = [] - totals: list[int] = [] + floors: list[tuple[int, int]] = [] for rank in range(len(layouts[0].attention_rows)): retained = workspace = 0 for ref, dispatched, layout in zip(refs, routed, layouts, strict=True): @@ -4554,13 +4567,9 @@ def _layout_checkpoint_floor( ) ) workspace = max(workspace, stage) - retained_by_rank.append(retained) - totals.append(retained + workspace) - retained = max(retained_by_rank) - workspace = max(totals) - retained - if moe: - workspace += self._te_workspace_growth_bytes() - return retained, workspace + floors.append((retained, workspace)) + growth = self._te_workspace_growth_bytes() if moe else 0 + return tuple((retained, workspace + growth) for retained, workspace in floors) def _layout_pricing_supported( self, topology: tuple[int, int, int, int], *, gradient_groups: bool @@ -4923,25 +4932,15 @@ def _checkpoint_layer_boundaries(self, retained: int) -> tuple[int, ...]: """Each decoder layer's saved-boundary bytes within the floor's retained term.""" return (retained // self._num_layers,) * self._num_layers - def _checkpoint_layer_boundary_sets( - self, - retained: int, - group_rows: tuple[tuple[int, bool], ...], - layouts: tuple[_GroupLayout, ...] | None, + def _layout_layer_boundaries( + self, layouts: tuple[_GroupLayout, ...] ) -> tuple[tuple[int, ...], ...]: - """Each rank's saved-boundary bytes per decoder layer, as the floor prices them. + """Each CP rank's saved-boundary bytes per decoder layer on its layouts. - On per-rank CP layouts (``_layout_checkpoint_floor``) a layer saves its + As ``_layout_checkpoint_rank_floors`` prices them: a layer saves its input in the GDN layout when it follows a GDN layer in its island and in - the attention layout otherwise, on that rank's rows. Elsewhere every - layer saves an equal share of the floor's ``retained``. + the attention layout otherwise, on that rank's rows. """ - if ( - layouts is None - or self._topology_key()[1] > 1 - or not all(grad for _, grad in group_rows) - ): - return (self._checkpoint_layer_boundaries(retained),) decoder = _language_model(self.runtime.model[0]).decoder gdn_inputs = [ getattr( @@ -4965,6 +4964,62 @@ def _checkpoint_layer_boundary_sets( ) return tuple(sets) + def _checkpoint_adapter_gradient_extra( + self, + slots: frozenset["LoRASlotRef"], + floor: tuple[int, int], + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, + routed_rows: tuple[int, ...] | None, + layouts: tuple[_GroupLayout, ...] | None, + ) -> int: + """Adapter gradients at the recompute backward's peak beyond ``floor``. + + ``floor`` is ``_checkpoint_memory_floor``'s (boundaries, workspace). + Every layer saves an equal share of its boundaries, except on per-rank + CP layouts (``_layout_checkpoint_floor``), where each rank releases its + own (``_layout_layer_boundaries``). A rank with fewer rows releases less + as backward proceeds, so its extra is larger, but its own floor is + smaller by what it never saved. Pair each rank's extra with its own + boundaries plus the larger of its workspace and the floor's, which + bounds that rank's floor, so a short CP2 sequence (all GDN rows on one + rank) does not add one rank's extra to the other's floor. + """ + retained, workspace = floor + if ( + layouts is None + or self._topology_key()[1] > 1 + or not all(grad for _, grad in group_rows) + ): + return self._checkpoint_adapter_gradient_bytes( + slots, self._checkpoint_layer_boundaries(retained) + ) + extras = [ + self._checkpoint_adapter_gradient_bytes(slots, boundaries) + for boundaries in self._layout_layer_boundaries(layouts) + ] + if not any(extras): + return 0 + floors = self._layout_checkpoint_rank_floors( + _language_model(self.runtime.model[0]).decoder.layers, + (None,) * len(group_rows) if slot_refs is None else slot_refs, + (None,) * len(group_rows) if routed_rows is None else routed_rows, + layouts, + ) + # No rank's pairing exceeds the floor, so every rank's own floor plus + # its extra fits within the floor plus this. + return max( + 0, + max( + rank_retained + max(rank_workspace, workspace) + extra + for (rank_retained, rank_workspace), extra in zip( + floors, extras, strict=True + ) + ) + - retained + - workspace, + ) + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -5036,13 +5091,14 @@ def _subforward_cost( # forward retention, including the cold fallback above. gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) gradient_slots = self._gradient_slots(group_rows, slot_refs) - # The busiest rank's extra, over each rank's own boundaries. adapter_gradient = ( - max( - self._checkpoint_adapter_gradient_bytes(gradient_slots, boundaries) - for boundaries in self._checkpoint_layer_boundary_sets( - checkpoint_retained, group_rows, group_layouts - ) + self._checkpoint_adapter_gradient_extra( + gradient_slots, + (checkpoint_retained, checkpoint_workspace), + group_rows, + slot_refs, + group_routed_rows, + group_layouts, ) if gradient else 0 @@ -8880,13 +8936,13 @@ def _estimate_required_memory_bytes_from_values( # as _subforward_cost. backward = self._checkpoint_input_gradient_bytes( group_rows, slot_refs - ) + max( - self._checkpoint_adapter_gradient_bytes( - self._gradient_slots(group_rows, slot_refs), boundaries - ) - for boundaries in self._checkpoint_layer_boundary_sets( - retained, group_rows, group_layouts - ) + ) + self._checkpoint_adapter_gradient_extra( + self._gradient_slots(group_rows, slot_refs), + (retained, workspace), + group_rows, + slot_refs, + group_routed_rows, + group_layouts, ) if profiled is None and any(grad for _, grad in group_rows): backward += _COLD_RECOMPUTE_TRANSIENT_BYTES diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index 38b52319e..a2e8914b4 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -304,22 +304,20 @@ def _plan_with(r, requests): return r._plan_flat_forward(requests) -def test_adapter_gradients_meet_each_ranks_own_layer_boundaries(monkeypatch): +def _pending_adapter_rank(monkeypatch, layout, rows, per_layer=300 * H): from art.megatron.lora import LoRASlotRef r = qwen36(rank()) - plan = _plan(r) - layout = _GroupLayout((400, 240), (160, 320), (0, 0)) monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) - monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((400, True),)) - monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (320,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((rows, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (rows,)) monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) - monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: 320) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: rows) policy = LoRASlotRef("checkpoint", "policy") monkeypatch.setattr( r, "_gradient_slots", lambda group_rows, slot_refs: frozenset({policy}) ) - pending = (300 * H,) * 40 + (0,) + pending = (per_layer,) * 40 + (0,) monkeypatch.setattr( r, "_pending_adapter_gradient_bytes", @@ -333,7 +331,8 @@ def test_adapter_gradients_meet_each_ranks_own_layer_boundaries(monkeypatch): def boundaries(rank_index): assert layout.gdn_rows is not None return [ - H * (layout.gdn_rows if is_gdn else layout.attention_rows)[rank_index] + H + * max(1, (layout.gdn_rows if is_gdn else layout.attention_rows)[rank_index]) for is_gdn in gdn_inputs ] @@ -346,13 +345,56 @@ def extra(saved): ), ) - # Each rank releases its own boundaries in layer order; the busiest extra wins. - expected = max(extra(boundaries(0)), extra(boundaries(1))) - assert r._plan_cost(plan).checkpoint_adapter_gradient == expected > 0 - # An even share of the busiest rank's boundaries would misplace the peak. - retained, _ = r._layout_checkpoint_floor(layers, (None,), (320,), (layout,)) - assert extra([retained // 40] * 40) != expected - assert r._checkpoint_layer_boundary_sets(retained, ((400, True),), (layout,)) == ( + floor = r._checkpoint_memory_floor( + ((rows, True),), None, routed_rows=(rows,), layouts=(layout,) + ) + floors = r._layout_checkpoint_rank_floors(layers, (None,), (rows,), (layout,)) + return r, boundaries, extra, floor, floors + + +def test_adapter_gradients_meet_each_ranks_own_layer_boundaries(monkeypatch): + layout = _GroupLayout((400, 240), (160, 320), (0, 0)) + r, boundaries, extra, (retained, workspace), floors = _pending_adapter_rank( + monkeypatch, layout, 400 + ) + assert r._layout_layer_boundaries((layout,)) == ( tuple(boundaries(0)), tuple(boundaries(1)), ) + # Every rank's boundaries sum to its own floor's. + assert [sum(boundaries(i)) for i in (0, 1)] == [floor[0] for floor in floors] + assert retained == max(floor[0] for floor in floors) + # Each rank releases its own boundaries in layer order; each extra sits on + # that rank's own floor, bounded by the floor's workspace. + expected = ( + max( + floor[0] + max(floor[1], workspace) + extra(boundaries(i)) + for i, floor in enumerate(floors) + ) + - retained + - workspace + ) + assert r._plan_cost(_plan(r)).checkpoint_adapter_gradient == expected > 0 + # An even share of the busiest rank's boundaries would misplace the peak. + assert extra([retained // 40] * 40) != expected + + +def test_a_rank_without_gdn_rows_pairs_its_extra_with_its_own_floor(monkeypatch): + # A single short sequence at CP2: every GDN row on rank 0 (Qwen3.6 2,047 + # tokens), with about 24 MB of expert LoRA gradients per layer. + layout = _GroupLayout((1536, 511), (2047, 0), (0, 0)) + r, boundaries, extra, (retained, workspace), floors = _pending_adapter_rank( + monkeypatch, layout, 2047, per_layer=6000 * H + ) + extras = [extra(boundaries(i)) for i in (0, 1)] + # Rank 1 releases almost nothing, so its extra is almost every gradient... + assert extras[1] > extras[0] and extras[1] > 6000 * H * 38 + cost = r._plan_cost(_plan(r)).checkpoint_adapter_gradient + # ...but it sits on rank 1's much smaller floor, not on rank 0's. + assert cost < extras[1] + # Every rank's own floor plus its own extra stays within the price. + for (rank_retained, rank_workspace), rank_extra in zip(floors, extras): + assert rank_retained + rank_workspace + rank_extra <= ( + retained + workspace + cost + ) + assert cost >= extras[0] From 22f67b0c407b3afd54f42c2d4e63f5182a22af31 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:13:33 +0000 Subject: [PATCH 13/19] Test split lower bounds and admission parity with pending gradients Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_layout_memory.py | 39 +++++++++++++------ 1 file changed, 28 insertions(+), 11 deletions(-) diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index 11fb7f992..fd09dbcbe 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -283,13 +283,18 @@ def test_softmax_offset_leaves_the_busiest_rank_floor(monkeypatch): assert r._plan_group_layouts(plan) is None +@pytest.mark.parametrize("per_layer", [0, 6000 * H]) @pytest.mark.parametrize( "lengths", [(2048, 1536, 1024, 512), (4099, 3, 5, 7), (1, 2, 3, 4, 5, 6, 7), (8191,)], ) -def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths): - # Even and skewed CP splits, odd row counts, and a single long sequence. +def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths, per_layer): + # Even and skewed CP splits, odd row counts, and a single long sequence; + # with pending policy gradients, even shares still price below each + # rank's paired extra. r = art_cp(rank(), monkeypatch) + if per_layer: + _with_policy_gradients(monkeypatch, r, (per_layer,) * 40 + (0,)) requests = _requests(lengths) plan = _plan_with(r, requests) lower = r._split_chunk_lower_cost( @@ -304,16 +309,10 @@ def _plan_with(r, requests): return r._plan_flat_forward(requests) -def _pending_adapter_rank(monkeypatch, layout, rows, per_layer=300 * H): +def _with_policy_gradients(monkeypatch, r, pending): + """Every gradient group trains the policy slot, with ``pending`` gradients.""" from art.megatron.lora import LoRASlotRef - r = qwen36(rank()) - monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) - monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((rows, True),)) - monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (rows,)) - monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) - monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: rows) - # The plan's gradient group trains the policy slot. policy = LoRASlotRef("checkpoint", "policy") groups = r._checkpoint_gradient_groups monkeypatch.setattr( @@ -323,12 +322,22 @@ def _pending_adapter_rank(monkeypatch, layout, rows, per_layer=300 * H): (policy, boundaries) for _, boundaries in groups(group_rows, slot_refs) ), ) - pending = (per_layer,) * 40 + (0,) monkeypatch.setattr( r, "_pending_adapter_gradient_bytes", lambda refs: pending if tuple(refs) else (), ) + + +def _pending_adapter_rank(monkeypatch, layout, rows, per_layer=300 * H): + r = qwen36(rank()) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((rows, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (rows,)) + monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: rows) + pending = (per_layer,) * 40 + (0,) + _with_policy_gradients(monkeypatch, r, pending) layers = r.runtime.model[0].decoder.layers gdn_inputs = [ layer._art_gdn_island_boundary.input_layout == "gdn" for layer in layers @@ -404,6 +413,14 @@ def test_a_rank_without_gdn_rows_pairs_its_extra_with_its_own_floor(monkeypatch) retained + workspace + cost ) assert cost >= extras[0] + # Admission prices the same extra: its gap to the plan cost stays the CP + # output coexistence alone (up to the safety factor's integer rounding). + plan = _plan(r) + gap = r._memory_check(plan).estimated_required_bytes - r._plan_cost(plan).required + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) + assert r._plan_cost(plan).checkpoint_adapter_gradient == 0 + plain = r._memory_check(plan).estimated_required_bytes - r._plan_cost(plan).required + assert abs(gap - plain) <= 1 def test_layout_gradient_groups_run_one_after_another_on_each_rank(monkeypatch): From 92dd4e0a5ad349e8480f891cc9626376bd91d7b1 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:33:34 +0000 Subject: [PATCH 14/19] Resolve each slot's pending gradients once across ranks Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 21 +++++++++++++++---- tests/unit/test_trainer_rank_layout_memory.py | 12 ++++++++--- 2 files changed, 26 insertions(+), 7 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index bf29dabff..73dc8a43e 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -4923,18 +4923,25 @@ def _checkpoint_gradient_groups( ) def _checkpoint_adapter_gradient_bytes( - self, groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]] + self, + groups: Sequence[tuple["LoRASlotRef | None", Sequence[int]]], + pending_by_slot: Mapping["LoRASlotRef", tuple[int, ...]] | None = None, ) -> int: """The recompute backward's adapter-gradient peak beyond released boundaries. ``groups`` gives each gradient group's adapter slot (None for the base model) and each decoder layer's saved-boundary bytes - (``_adapter_gradient_walk``). + (``_adapter_gradient_walk``). ``pending_by_slot`` reuses slots' pending + gradients already resolved (``_pending_adapter_gradient_bytes``). """ chains = [] for slot, boundaries in groups: pending = ( - () if slot is None else self._pending_adapter_gradient_bytes((slot,)) + () + if slot is None + else pending_by_slot[slot] + if pending_by_slot is not None + else self._pending_adapter_gradient_bytes((slot,)) ) if pending and len(pending) != len(boundaries) + 1: return 0 @@ -5046,9 +5053,15 @@ def _checkpoint_adapter_gradient_extra( ): return self._checkpoint_adapter_gradient_bytes(groups) slots = [slot for slot, _ in groups] + # Every rank walks the same slots' gradients: resolve them once. + pending = { + slot: self._pending_adapter_gradient_bytes((slot,)) + for slot in slots + if slot is not None + } extras = [ self._checkpoint_adapter_gradient_bytes( - tuple(zip(slots, boundaries, strict=True)) + tuple(zip(slots, boundaries, strict=True)), pending ) for boundaries in self._layout_layer_boundaries(layouts) ] diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index fd09dbcbe..df889e67a 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -450,9 +450,13 @@ def test_layout_gradient_groups_run_one_after_another_on_each_rank(monkeypatch): (slots[0],): (6000 * H,) * 40 + (0,), (slots[1],): (100 * H,) * 40 + (5 * H,), } - monkeypatch.setattr( - r, "_pending_adapter_gradient_bytes", lambda refs: pending.get(tuple(refs), ()) - ) + walks = [] + + def pending_gradients(refs): + walks.append(tuple(refs)) + return pending.get(tuple(refs), ()) + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending_gradients) floor = r._checkpoint_memory_floor( group_rows, None, routed_rows=routed, layouts=layouts ) @@ -484,3 +488,5 @@ def test_layout_gradient_groups_run_one_after_another_on_each_rank(monkeypatch): == expected > 0 ) + # Each slot's module walk runs once, not once per rank. + assert walks == [(slot,) for slot in slots] From bd5a83a355910726d27744aff75b80b78db28584 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:35:05 +0000 Subject: [PATCH 15/19] Test the split lower bound with two gradient slots Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_layout_memory.py | 46 +++++++++++++++++++ 1 file changed, 46 insertions(+) diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index df889e67a..1cb4d39e2 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -305,6 +305,52 @@ def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths, per assert r._layout_pricing_supported((1, 1, 2, 1), gradient_groups=True) +@pytest.mark.parametrize( + "lengths", [(2048, 1536, 1024, 512), (4099, 3, 5, 7), (8191, 64)] +) +def test_split_lower_bound_stays_below_two_slots_layout_cost(monkeypatch, lengths): + # Two gradient groups of different slots (the first request alone, then + # the rest), each with pending gradients, one outweighing its boundaries. + from art.megatron.lora import LoRASlotRef + + r = art_cp(rank(), monkeypatch) + base = LoRASlotRef("checkpoint", None) + monkeypatch.setattr( + r, + "_group_active_request_indices", + lambda requests, **_: ( + ((None, True), (0,)), + ((base, True), tuple(range(1, len(requests)))), + ), + ) + slots = (LoRASlotRef("checkpoint", "policy"), LoRASlotRef("checkpoint", "other")) + groups = r._checkpoint_gradient_groups + monkeypatch.setattr( + r, + "_checkpoint_gradient_groups", + lambda group_rows, slot_refs: tuple( + (slot, boundaries) + for slot, (_, boundaries) in zip(slots, groups(group_rows, slot_refs)) + ), + ) + pending = { + (slots[0],): (6000 * H,) * 40 + (0,), + (slots[1],): (300 * H,) * 40 + (7 * H,), + } + monkeypatch.setattr( + r, "_pending_adapter_gradient_bytes", lambda refs: pending.get(tuple(refs), ()) + ) + requests = _requests(lengths) + plan = _plan_with(r, requests) + assert len(plan.groups) == 2 and r._plan_group_layouts(plan) is not None + exact = r._plan_cost(plan) + assert exact.checkpoint_adapter_gradient > 0 + lower = r._split_chunk_lower_cost( + requests, tuple(item.input_tokens for item in requests), checkpoint=Unset + ) + assert lower.required <= exact.required + + def _plan_with(r, requests): return r._plan_flat_forward(requests) From 632dea956467e60b08c406a1f7bb5d1af07bb142 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Mon, 28 Sep 2026 08:43:27 +0000 Subject: [PATCH 16/19] Pin the two-slot lower bound's own extra Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_layout_memory.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py index 1cb4d39e2..326da5caf 100644 --- a/tests/unit/test_trainer_rank_layout_memory.py +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -310,7 +310,7 @@ def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths, per ) def test_split_lower_bound_stays_below_two_slots_layout_cost(monkeypatch, lengths): # Two gradient groups of different slots (the first request alone, then - # the rest), each with pending gradients, one outweighing its boundaries. + # the rest), each with pending gradients. from art.megatron.lora import LoRASlotRef r = art_cp(rank(), monkeypatch) @@ -348,6 +348,8 @@ def test_split_lower_bound_stays_below_two_slots_layout_cost(monkeypatch, length lower = r._split_chunk_lower_cost( requests, tuple(item.input_tokens for item in requests), checkpoint=Unset ) + # The even-share lower bound prices both slots' extra too, and stays below. + assert lower.checkpoint_adapter_gradient > 0 assert lower.required <= exact.required From 653136d5e85ebd14d0cbd3b0b5b9d400dc3239b0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 05:59:37 +0000 Subject: [PATCH 17/19] Price CP2 recompute per rank on each layer type's layout, on the extracted layout Port #978 (through 6dc567b03) onto the extracted planner modules: every CP2 rank's checkpoint boundaries and recomputed layer are priced on its own attention and GDN layouts plus the CP executor's retained stage records, with adapter gradients paired per rank. Grouped planner-miss replay freezes each group's layouts and the decoder's per-layer input layouts (facts version 4). Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 6 +- src/art/megatron/context_parallel/executor.py | 107 ++++ src/art/megatron/context_parallel/runtime.py | 49 ++ src/art/trainer_rank/_impl.py | 46 +- src/art/trainer_rank/_memory.py | 366 ++++++++++- src/art/trainer_rank/_micro_batch_planner.py | 94 +++ src/art/trainer_rank/_planner_misses.py | 1 + src/art/trainer_rank/_planner_replay.py | 99 ++- .../test_context_parallel_retained_bytes.py | 157 +++++ tests/unit/test_grouped_planner_replay.py | 102 +++ .../test_trainer_rank_admission_inputs.py | 1 + .../test_trainer_rank_checkpoint_memory.py | 9 +- tests/unit/test_trainer_rank_layout_memory.py | 595 ++++++++++++++++++ tests/unit/test_trainer_rank_moe_memory.py | 57 +- 14 files changed, 1650 insertions(+), 39 deletions(-) create mode 100644 tests/unit/test_context_parallel_retained_bytes.py create mode 100644 tests/unit/test_trainer_rank_layout_memory.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index f7360485f..f4506e8af 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -245,6 +245,8 @@ jobs: tests/unit/test_trainer_rank_pending_memory.py \ tests/unit/test_trainer_rank_shared_memory.py \ tests/unit/test_trainer_rank_converted_memory.py \ + tests/unit/test_trainer_rank_layout_memory.py \ + tests/unit/test_context_parallel_retained_bytes.py \ tests/unit/test_trainer_rank_split.py \ tests/unit/test_megatron_compile_garbage.py \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ @@ -295,4 +297,6 @@ jobs: --ignore=tests/unit/test_trainer_rank_pending_memory.py \ --ignore=tests/unit/test_trainer_rank_shared_memory.py \ --ignore=tests/unit/test_megatron_compile_garbage.py \ - --ignore=tests/unit/test_trainer_rank_converted_memory.py + --ignore=tests/unit/test_trainer_rank_converted_memory.py \ + --ignore=tests/unit/test_trainer_rank_layout_memory.py \ + --ignore=tests/unit/test_context_parallel_retained_bytes.py diff --git a/src/art/megatron/context_parallel/executor.py b/src/art/megatron/context_parallel/executor.py index 4015011fe..c59751966 100644 --- a/src/art/megatron/context_parallel/executor.py +++ b/src/art/megatron/context_parallel/executor.py @@ -35,6 +35,7 @@ DkvReducePlan, ExactMaskMetadata, FlexMaskSpec, + RankRuntimePlan, StageExecutionSpec, StagePlan, TokenRange, @@ -1550,6 +1551,112 @@ def _merge_stage_output_grads_from_tape( return stage_out_grads, stage_lse_grads +def minimum_retained_bytes_per_row( + *, + q_heads: int, + kv_heads: int, + head_dim: int, + value_head_dim: int, + element_size: int, +) -> int: + """The least ``retained_stage_record_bytes`` keeps per own row. + + Every rank with rows runs a local stage over all of them (each row attends + to itself); an aligned one keeps at least the contiguous copies of + multi-head views, flex's output and its two LSEs. + """ + copies = (q_heads > 1) * q_heads * head_dim + (kv_heads > 1) * kv_heads * ( + head_dim + value_head_dim + ) + return (copies + q_heads * value_head_dim) * element_size + 2 * q_heads * 4 + + +def retained_stage_record_bytes( + rank_plan: RankRuntimePlan, + *, + q_heads: int, + kv_heads: int, + head_dim: int, + value_head_dim: int, + element_size: int, + block_size: SparseBlockSize, +) -> int: + """Bytes a recomputed attention keeps for backward beyond its own-row tensors. + + A size-only mirror of ``_forward_stage_records`` with ``record_for_backward`` + and of ``_run_stage_attention``; keep the three in step. Every stage that + runs keeps the Q/K/V flex consumed: a padded copy when the execution length + differs, else a contiguous copy of the permuted ``q_flat``/``k_flat`` view; + partial-range gathers and remote fetch buffers are kept as the stage's + inputs. It keeps flex's output and LSE at the execution length, and + logical-length copies of them when padded, plus flex's own LSE beside the + normalized one the FLASH backend returns (counted on every backend). Every + producing stage after the first keeps a merge-tape clone of the + accumulators. Accumulators themselves are transient. + """ + own = int(rank_plan.local_valid_lengths[0]) if rank_plan.local_valid_lengths else 0 + accum_size = 4 if element_size < 4 else element_size + q_row = q_heads * head_dim * element_size + k_row = kv_heads * head_dim * element_size + v_row = kv_heads * value_head_dim * element_size + out_row = q_heads * value_head_dim * element_size + lse_row = q_heads * 4 + tape_row = q_heads * (value_head_dim + 1) * accum_size + total = 0 + local_produced = False + tapes: list[int] = [] + for stage in _ordered_stage_plans(rank_plan.stage_plans): + if not (stage.q_len > 0 and stage.k_len > 0 and stage.slices): + continue + q_len = _logical_stage_q_len(stage) + k_len = _logical_stage_k_len(stage) + q_pad, k_pad, _family = select_sparse_execution_family( + is_local_stage=bool(stage.is_local_stage), + q_len=int(stage.q_len), + k_len=int(stage.k_len), + block_size=block_size, + ) + q_full = _ranges_cover_full_length(stage.owner_local_q_ranges, length=own) + # Queries: the full range is a permuted view of q_flat; a partial range + # is gathered into a contiguous tensor kept as the stage's input. + if not q_full: + total += q_row * q_len + if q_pad != q_len: + total += q_row * q_pad + elif q_full: + # A view of the projection output: flex copies it unless it is + # already contiguous, which one head alone does not guarantee + # (a fused QKV split keeps the projection's token stride). + total += q_row * q_len + # Keys and values: local ranges as for queries; remote ones land in + # contiguous head-major fetch buffers kept as the stage's inputs. + k_full = bool(stage.is_local_stage) and _ranges_cover_full_length( + stage.owner_local_k_ranges, length=own + ) + if not k_full: + total += (k_row + v_row) * k_len + if k_pad != k_len: + total += (k_row + v_row) * k_pad + elif k_full: + total += (k_row + v_row) * k_len + total += (out_row + 2 * lse_row) * q_pad + if q_pad != q_len: + total += (out_row + lse_row) * q_len + tape = tape_row * (own if q_full else q_len) + if not q_full: + # Its int64 row index is kept even by the first producing stage. + total += 8 * q_len + if stage.is_local_stage: + local_produced = True + else: + tapes.append(tape) + # The first producing stage records no tape: the local stage when it ran, + # else whichever remote stage is ready first, so drop the smallest tape. + if tapes and not local_produced: + tapes.remove(min(tapes)) + return total + sum(tapes) + + def _forward_stage_records( *, q_flat: torch.Tensor, diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 531cb8829..2f7ac2dad 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -433,6 +433,55 @@ def context_parallel_rank_model_token_counts( ) +def context_parallel_rank_layouts( + *, + group_ids: torch.Tensor, + parent_ids: torch.Tensor, + topology: ParallelTopology, + config: ContextParallelConfig, + original_seq_len: int, + build_gdn_execution_spec: bool, + gdn_planner_config: Any | None = None, +) -> tuple[tuple[int, ...], tuple[int, ...] | None, tuple[RankRuntimePlan, ...]]: + """Each CP rank's attention rows, GDN rows and attention stage plan. + + Uses the cached planning bundle and per-rank runtime plans that execution + builds, so a memory estimate sees the layouts the ranks will run. + """ + planning_key, bundle, _group_ids_cpu, _parent_ids_cpu = ( + _get_or_build_planning_bundle( + group_ids=group_ids, + parent_ids=parent_ids, + topology=topology, + config=config, + original_seq_len=original_seq_len, + build_gdn_execution_spec=build_gdn_execution_spec, + ) + ) + attention = tuple(bundle.token_layout_index.token_counts_by_rank) + gdn = None + if build_gdn_execution_spec: + gdn = tuple( + _plan_gdn_global_execution( + planning_key=planning_key, + bundle=bundle, + topology=topology, + gdn_planner_config=gdn_planner_config, + ).gdn_token_counts_by_rank + ) + plans = tuple( + _get_or_build_bundle_rank_plan( + planning_key=planning_key, + bundle=bundle, + original_seq_len=original_seq_len, + target_rank=rank, + block_size=config.block_size, + ) + for rank in range(len(attention)) + ) + return attention, gdn, plans + + def context_parallel_model_token_total( *, group_ids: torch.Tensor, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 2ec11b766..7aef00465 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1069,6 +1069,17 @@ def signature(self) -> _MemorySignature: _AnyForwardPlan = _FlatForwardPlan | _SplitForwardPlan +@dataclass(frozen=True) +class _GroupLayout: + """One packed group's CP layouts on every rank, for layout-aware pricing.""" + + attention_rows: tuple[int, ...] + gdn_rows: tuple[int, ...] | None + # What each rank's recomputed CP attention keeps for backward beyond its + # own-row activations (``retained_stage_record_bytes``). + attention_retained: tuple[int, ...] + + @dataclass(frozen=True) class _SubforwardCost: """Memory terms of one candidate subforward while all graphs stay live. @@ -3041,6 +3052,7 @@ def _subforward_cost( gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), group_routed_rows: tuple[int, ...] | None = None, + group_layouts: tuple[_GroupLayout, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, head_backward_traced: bool = False, @@ -3049,7 +3061,11 @@ def _subforward_cost( hybridep_growth_bytes: int = 0, ) -> _SubforwardCost: checkpoint_memory = self._checkpoint_memory_floor( - group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows + group_rows, + slot_refs, + gdn_segments, + routed_rows=group_routed_rows, + layouts=group_layouts, ) required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, @@ -3059,6 +3075,7 @@ def _subforward_cost( gdn_segments=gdn_segments, group_rows=group_rows, group_routed_rows=group_routed_rows, + group_layouts=group_layouts, slot_refs=slot_refs, head_workspace_bytes=head_workspace_bytes, checkpoint_floor=checkpoint_floor, @@ -3080,13 +3097,21 @@ def _subforward_cost( ) # Input gradients live at the recomputed layer's peak; kept out of # forward retention, including the cold fallback above. + # Per-rank layouts price each rank's own boundaries; the input-gradient + # allowance keeps the busiest rank's (``_checkpoint_input_gradient_bytes``). gradient = self._checkpoint_input_gradient_bytes( - group_rows, slot_refs, retained=checkpoint_retained + group_rows, + slot_refs, + retained=checkpoint_retained if group_layouts is None else None, ) gradient_slots = self._gradient_slots(group_rows, slot_refs) adapter_gradient = ( - self._checkpoint_adapter_gradient_bytes( - self._checkpoint_gradient_groups(group_rows, slot_refs) + self._checkpoint_adapter_gradient_extra( + (checkpoint_retained, checkpoint_workspace), + group_rows, + slot_refs, + group_routed_rows, + group_layouts, ) if gradient else 0 @@ -3099,7 +3124,7 @@ def _subforward_cost( forward_required = required if gradient: head_stage = self._checkpoint_head_stage_bytes( - head_workspace_bytes, gradient, group_rows, slot_refs + head_workspace_bytes, gradient, group_rows, slot_refs, group_layouts ) peak = checkpoint_workspace + adapter_gradient if head_stage is not None: @@ -4929,6 +4954,17 @@ def _cp_group_model_tokens( ) _group_head_workspace_bytes = _memory._group_head_workspace_bytes + _layer_gdn_inputs = _memory._layer_gdn_inputs + _layout_checkpoint_floor = _memory._layout_checkpoint_floor + _layout_checkpoint_rank_floors = _memory._layout_checkpoint_rank_floors + _layout_pricing_supported = _memory._layout_pricing_supported + _minimum_layouts = _memory._minimum_layouts + _layout_layer_boundaries = _memory._layout_layer_boundaries + _layout_adapter_gradient_bytes = _memory._layout_adapter_gradient_bytes + _checkpoint_adapter_gradient_extra = _memory._checkpoint_adapter_gradient_extra + _recomputed_mixer_widths = _memory._recomputed_mixer_widths + _plan_group_layouts = _micro_batch_planner._plan_group_layouts + _compute_group_layouts = _micro_batch_planner._compute_group_layouts _triton_min_rows = _memory._triton_min_rows _head_backward_traced = _memory._head_backward_traced _te_workspace_growth_bytes = _memory._te_workspace_growth_bytes diff --git a/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index 6a4654222..b0d0da39b 100644 --- a/src/art/trainer_rank/_memory.py +++ b/src/art/trainer_rank/_memory.py @@ -13,7 +13,7 @@ from __future__ import annotations -from collections.abc import Iterable, Iterator, Sequence +from collections.abc import Iterable, Iterator, Mapping, Sequence from contextlib import nullcontext import hashlib import math @@ -27,7 +27,12 @@ import torch.distributed as dist from art.megatron.lora import LoRASlotRef - from art.trainer_rank._impl import AdapterSelection, AnyForwardInput, TrainerRank + from art.trainer_rank._impl import ( + AdapterSelection, + AnyForwardInput, + TrainerRank, + _GroupLayout, + ) def _split_required_memory(costs: Sequence[_impl._SubforwardCost]) -> int: @@ -501,8 +506,10 @@ def _moe_workspace_from_terms( """``routed_rows`` is what this rank's experts receive at balanced routing (``rows`` by default); the shared expert's part stays on the local rows.""" coefficient, stages, shared = terms - routed = rows if routed_rows is None else max(0, min(rows, routed_rows)) - return (rows - routed) * shared + ( + # A rank can receive more routed rows than it holds (an uneven CP + # split); the shared part then stays on the routed rows' count. + routed = rows if routed_rows is None else max(0, routed_rows) + return max(0, rows - routed) * shared + ( max( routed * coefficient, *(routed * per_row + fixed for per_row, fixed in stages), @@ -608,19 +615,22 @@ def _checkpoint_memory_floor( slot_refs: tuple["LoRASlotRef | None", ...] | None = None, gdn_segments: int = 0, routed_rows: tuple[int, ...] | None = None, + layouts: tuple[_GroupLayout, ...] | None = None, ) -> tuple[int, int]: """Conservative saved-boundary charge and one recomputed layer's workspace. ``routed_rows`` are each group's balanced dispatched rows per rank - (``_plan_group_routed_rows``); by default, its local rows. Eligibility is - ``_checkpoint_layers``; the arithmetic, shared with grouped CPU replay, is - ``_checkpoint_floor_from_facts``. + (``_plan_group_routed_rows``); by default, its local rows. With + ``layouts`` (``_plan_group_layouts``), price every rank on its own CP + layouts instead of the busiest rank's rows (``_layout_checkpoint_floor``). + Eligibility is ``_checkpoint_layers``; the arithmetic, shared with grouped + CPU replay, is ``_checkpoint_floor_from_facts``. """ layers = _checkpoint_layers(self, group_rows) if not layers: return 0, 0 retained, workspace = _checkpoint_floor_from_facts( - self, group_rows, slot_refs, gdn_segments, layers, routed_rows + self, group_rows, slot_refs, gdn_segments, layers, routed_rows, layouts ) # HybridEP runtime state is intentionally outside grouped CPU replay. if ( @@ -652,6 +662,7 @@ def _checkpoint_floor_from_facts( gdn_segments: int, layers: int, routed_rows: tuple[int, ...] | None = None, + layouts: tuple[_GroupLayout, ...] | None = None, ) -> tuple[int, int]: """Count actual local full/uniform/1 boundaries, including aliases, rather than claiming measured distinct storage. Only this call's new groups @@ -696,8 +707,10 @@ def _checkpoint_floor_from_facts( roots = gdn_segments + (tp - 1) * sum(grad for _, grad in group_rows) workspace += math.ceil(roots * self._gdn_segment_layer_bytes()) return retained, workspace - # Boundaries on the busiest rank's rows and one recomputed layer. routed = (None,) * len(group_rows) if routed_rows is None else routed_rows + if layouts is not None and all(grad for _, grad in group_rows): + return self._layout_checkpoint_floor(refs, routed, layouts) + # Boundaries on the busiest rank's rows and one recomputed layer. retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 # Beside the mixer, the recomputed layer keeps its post-mixer residual @@ -846,6 +859,7 @@ def _checkpoint_gradient_groups( def _checkpoint_adapter_gradient_bytes( self: TrainerRank, groups: Sequence[tuple[LoRASlotRef | None, Sequence[int]]], + pending_by_slot: Mapping[LoRASlotRef, tuple[int, ...]] | None = None, *, head: bool = False, ) -> int: @@ -853,12 +867,20 @@ def _checkpoint_adapter_gradient_bytes( ``groups`` gives each gradient group's adapter slot (None for the base model) and each decoder layer's saved-boundary bytes - (``_adapter_gradient_walk``). With ``head``, those live while a group's - head runs its backward instead (``_adapter_gradient_head``). + (``_adapter_gradient_walk``). ``pending_by_slot`` reuses slots' pending + gradients already resolved (``_pending_adapter_gradient_bytes``). With + ``head``, those live while a group's head runs its backward instead + (``_adapter_gradient_head``). """ chains = [] for slot, boundaries in groups: - pending = () if slot is None else self._pending_adapter_gradient_bytes((slot,)) + pending = ( + () + if slot is None + else pending_by_slot[slot] + if pending_by_slot is not None + else self._pending_adapter_gradient_bytes((slot,)) + ) if pending and len(pending) != len(boundaries) + 1: return 0 chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) @@ -1078,6 +1100,7 @@ def _checkpoint_head_stage_bytes( gradient: int, group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None, + layouts: tuple[_GroupLayout, ...] | None = None, ) -> int | None: """The head backward's peak beyond the boundaries and ``gradient``. @@ -1095,12 +1118,17 @@ def _checkpoint_head_stage_bytes( shares the decoder stage. That bounds heads whose backward is the traced one (``_head_backward_traced``); callers keep the unstaged price beside any other, whose buffers can exceed ``head_workspace_bytes``. + On per-rank CP layouts, each rank's boundaries are within the floor's, + so the largest rank's adapter term bounds each. """ if not head_workspace_bytes or not self._checkpoint_gradient_covered( group_rows, slot_refs ): return None rows = sum(rows for rows, grad in group_rows if grad) + adapters = self._layout_adapter_gradient_bytes( + group_rows, slot_refs, layouts, head=True + ) # A rank's chunk below the fused minimum (CP splits the projected rows # unevenly) takes the FP32 fallback. fallback = _impl._HEAD_FALLBACK_BUFFERS * self._head_workspace_bytes( @@ -1114,8 +1142,12 @@ def _checkpoint_head_stage_bytes( + 2 * gradient + rows * self._backward_row_state_bytes() + self._te_workspace_growth_bytes() - + self._checkpoint_adapter_gradient_bytes( - self._checkpoint_gradient_groups(group_rows, slot_refs), head=True + + ( + self._checkpoint_adapter_gradient_bytes( + self._checkpoint_gradient_groups(group_rows, slot_refs), head=True + ) + if adapters is None + else max(adapters, default=0) ) ) @@ -1191,18 +1223,296 @@ def _recomputed_mixer_bytes_per_token(self: TrainerRank) -> int: more rows; its hidden-width input exchange and value-width output allowance price that (88 KB measured at CP2, 94 KB priced). """ + return max(self._recomputed_mixer_widths().values(), default=0) + + +def _layer_gdn_inputs(self: TrainerRank) -> tuple[bool, ...]: + """Per decoder layer, whether its saved input arrives in the GDN layout. + + It does when the layer follows a GDN layer in its island + (``_art_gdn_island_boundary``); otherwise it is in the attention layout. + """ + decoder = _impl._language_model(self.runtime.model[0]).decoder + return tuple( + getattr(getattr(layer, "_art_gdn_island_boundary", None), "input_layout", "") + == "gdn" + for layer in decoder.layers + ) + + +def _layout_checkpoint_floor( + self: TrainerRank, + refs: tuple["LoRASlotRef | None", ...], + routed: tuple[int | None, ...], + layouts: tuple[_GroupLayout, ...], +) -> tuple[int, int]: + """The largest rank's boundaries and the rest of the largest rank total. + + The two sum to the largest of ``_layout_checkpoint_rank_floors``' totals. + """ + floors = self._layout_checkpoint_rank_floors(refs, routed, layouts) + retained = max(retained for retained, _ in floors) + return retained, max(map(sum, floors)) - retained + + +def _layout_checkpoint_rank_floors( + self: TrainerRank, + refs: tuple["LoRASlotRef | None", ...], + routed: tuple[int | None, ...], + layouts: tuple[_GroupLayout, ...], +) -> tuple[tuple[int, int], ...]: + """Each CP rank's boundaries and recomputed layer on its own layouts. + + A saved layer input arrives in the GDN layout when the layer follows a + GDN layer in its island (``_art_gdn_island_boundary``), and in the + attention layout otherwise. A recomputed attention layer keeps its + activations on its attention rows plus what the CP executor keeps for + backward; a GDN layer keeps the GDN width on its GDN rows. Either one's + residual, norm, routing state and MoE stage use that layer's rows, and + routed rows the EP share. Traced CP2 ranks: boundaries match exactly and + attention within 0.3%. Returns each rank's (boundaries, workspace). + """ + hidden = self._hidden_size * 2 + inputs = self._layer_gdn_inputs() + gdn_inputs = sum(inputs) + attention_inputs = len(inputs) - gdn_inputs + widths = self._recomputed_mixer_widths(stage_buffers=False) + moe = self._checkpoint_moe_bytes_per_token() + beside = 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe else 0 + floors: list[tuple[int, int]] = [] + for rank in range(len(layouts[0].attention_rows)): + retained = workspace = 0 + for ref, dispatched, layout in zip(refs, routed, layouts, strict=True): + attention = max(1, layout.attention_rows[rank]) + gdn = ( + attention if layout.gdn_rows is None else max(1, layout.gdn_rows[rank]) + ) + retained += hidden * (attention_inputs * attention + gdn_inputs * gdn) + for kind, rows, extra in ( + ("attention", attention, layout.attention_retained[rank]), + ("gdn", gdn, 0), + ): + if kind not in widths: + continue + stage = ( + rows * (widths[kind] + beside) + + extra + + self._moe_workspace_bytes( + rows, + routed_rows=dispatched, + checkpoint_grad=True, + slot_ref=ref, + ) + ) + workspace = max(workspace, stage) + floors.append((retained, workspace)) + growth = self._te_workspace_growth_bytes() if moe else 0 + return tuple((retained, workspace + growth) for retained, workspace in floors) + + +def _layout_pricing_supported( + self: TrainerRank, topology: tuple[int, int, int, int], *, gradient_groups: bool +) -> bool: + """Whether layout-aware pricing models this runtime and plan shape.""" + _dp, tp, cp, pp = topology + if (tp, cp, pp) != (1, 2, 1) or not gradient_groups: + return False + geometry = self._geometry + if not geometry.num_attention_heads or not geometry.kv_channels: + return False + if len(self.runtime.model) != 1: + return False + try: + decoder = _impl._language_model(self.runtime.model[0]).decoder + from art.megatron.context_parallel.core_attention import ( + ArtContextParallelCoreAttention, + ) + except (AttributeError, RuntimeError, ModuleNotFoundError): + return False + layers = getattr(decoder, "layers", None) + if layers is None: + return False + for layer in layers: + boundary = getattr(layer, "_art_gdn_island_boundary", None) + if boundary is not None and boundary.is_gdn: + continue + if self._gdn_layers and boundary is None: + return False + core = getattr(getattr(layer, "self_attention", None), "core_attention", None) + if ( + type(core) is not ArtContextParallelCoreAttention + or getattr(core, "softmax_offset", None) is not None + ): + return False + return True + + +def _minimum_layouts( + self: TrainerRank, physical_rows: Sequence[int], cp: int +) -> tuple[_GroupLayout, ...]: + """Even-share layouts keeping the least attention state: a lower bound. + + Every rank's total grows with its own rows, and every rank keeps at + least an aligned local stage's state per row (each row attends to + itself), so the largest rank's total is at least the total at the mean + rows. The mean is at least the floor of an even share, which is why + this rounds down; rounding up can exceed a split's exact cost. + """ + from art.megatron.context_parallel.executor import ( + minimum_retained_bytes_per_row, + ) + + geometry = self._geometry + per_row = minimum_retained_bytes_per_row( + q_heads=int(geometry.num_attention_heads), + kv_heads=int(geometry.num_query_groups), + head_dim=int(geometry.kv_channels), + value_head_dim=int(geometry.kv_channels), + element_size=self._param_dtype_size, + ) + return tuple( + _impl._GroupLayout( + attention_rows=(rows // cp,) * cp, + gdn_rows=(rows // cp,) * cp if self._gdn_layers else None, + attention_retained=(rows // cp * per_row,) * cp, + ) + for rows in physical_rows + ) + + +def _layout_layer_boundaries( + self: TrainerRank, layouts: tuple[_GroupLayout, ...] +) -> tuple[tuple[tuple[int, ...], ...], ...]: + """Each CP rank's saved-boundary bytes per group and decoder layer. + + As ``_layout_checkpoint_rank_floors`` prices them: a layer saves its + input in the GDN layout when it follows a GDN layer in its island and in + the attention layout otherwise, on that rank's rows of the group. + """ + gdn_inputs = self._layer_gdn_inputs() + hidden = self._hidden_size * 2 + ranks = [] + for rank in range(len(layouts[0].attention_rows)): + groups = [] + for layout in layouts: + attention = max(1, layout.attention_rows[rank]) + gdn = ( + attention if layout.gdn_rows is None else max(1, layout.gdn_rows[rank]) + ) + groups.append( + tuple(hidden * (gdn if is_gdn else attention) for is_gdn in gdn_inputs) + ) + ranks.append(tuple(groups)) + return tuple(ranks) + + +def _layout_adapter_gradient_bytes( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, + layouts: tuple[_GroupLayout, ...] | None, + *, + head: bool = False, +) -> list[int] | None: + """Each CP rank's ``_checkpoint_adapter_gradient_bytes`` on its own rows. + + Over that rank's boundaries (``_layout_layer_boundaries``). None + without per-rank layouts: every layer then saves each gradient group's + rows (``_checkpoint_gradient_groups``). + """ + if ( + layouts is None + or self._topology_key()[1] > 1 + or not all(grad for _, grad in group_rows) + ): + return None + slots = [ + slot for slot, _ in self._checkpoint_gradient_groups(group_rows, slot_refs) + ] + # Every rank walks the same slots' gradients: resolve them once. + pending = { + slot: self._pending_adapter_gradient_bytes((slot,)) + for slot in slots + if slot is not None + } + return [ + self._checkpoint_adapter_gradient_bytes( + tuple(zip(slots, boundaries, strict=True)), pending, head=head + ) + for boundaries in self._layout_layer_boundaries(layouts) + ] + + +def _checkpoint_adapter_gradient_extra( + self: TrainerRank, + floor: tuple[int, int], + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, + routed_rows: tuple[int, ...] | None, + layouts: tuple[_GroupLayout, ...] | None, +) -> int: + """Adapter gradients at the recompute backward's peak beyond ``floor``. + + ``floor`` is ``_checkpoint_memory_floor``'s (boundaries, workspace). + Every layer saves each gradient group's rows + (``_checkpoint_gradient_groups``), except on per-rank CP layouts + (``_layout_checkpoint_floor``), where each rank releases its own + (``_layout_layer_boundaries``). A rank with fewer rows releases less + as backward proceeds, so its extra is larger, but its own floor is + smaller by what it never saved. Pair each rank's extra with its own + boundaries plus the larger of its workspace and the floor's, which + bounds that rank's floor, so a short CP2 sequence (all GDN rows on one + rank) does not add one rank's extra to the other's floor. + """ + retained, workspace = floor + extras = self._layout_adapter_gradient_bytes(group_rows, slot_refs, layouts) + if extras is None: + return self._checkpoint_adapter_gradient_bytes( + self._checkpoint_gradient_groups(group_rows, slot_refs) + ) + if not any(extras): + return 0 + assert layouts is not None # extras come only from per-rank layouts + floors = self._layout_checkpoint_rank_floors( + (None,) * len(group_rows) if slot_refs is None else slot_refs, + (None,) * len(group_rows) if routed_rows is None else routed_rows, + layouts, + ) + # No rank's pairing exceeds the floor, so every rank's own floor plus + # its extra fits within the floor plus this. + return max( + 0, + max( + rank_retained + max(rank_workspace, workspace) + extra + for (rank_retained, rank_workspace), extra in zip( + floors, extras, strict=True + ) + ) + - retained + - workspace, + ) + + +def _recomputed_mixer_widths( + self: TrainerRank, *, stage_buffers: bool = True +) -> dict[str, int]: + """Recomputed attention and GDN mixer bytes per row, by layer type. + + ``stage_buffers=False`` leaves out the CP attention stage allowance, for + callers that price the executor's stage buffers from its stage plan. + """ geometry = self._geometry hidden = self._hidden_size tp = max(1, self._topology_key()[1]) cp = self._topology_key()[2] > 1 - widths = [] + widths: dict[str, float] = {} if self._gdn_layers < self._num_layers: attention, _gdn = self._mixer_activation_widths() - if cp: + if cp and stage_buffers: q = geometry.num_attention_heads * geometry.kv_channels or hidden kv = geometry.num_query_groups * geometry.kv_channels or hidden attention += (3 * q + 2 * kv) / tp - widths.append(attention) + widths["attention"] = attention if self._gdn_layers: key = geometry.gdn_key_heads * geometry.gdn_key_head_dim value = geometry.gdn_value_heads * geometry.gdn_value_head_dim @@ -1211,8 +1521,8 @@ def _recomputed_mixer_bytes_per_token(self: TrainerRank) -> int: gdn = hidden + (2 * key + normalized + 6 * value + chunk) / tp if cp: gdn += hidden + value / tp - widths.append(gdn) - return int(max(widths, default=0) * self._param_dtype_size) + widths["gdn"] = gdn + return {kind: int(width * self._param_dtype_size) for kind, width in widths.items()} def _adapter_gradient_head( @@ -1584,6 +1894,7 @@ def _memory_check( gdn_segments=forward.grad_segment_count, group_rows=self._plan_group_rows(forward), group_routed_rows=self._plan_group_routed_rows(forward), + group_layouts=self._plan_group_layouts(forward), slot_refs=tuple(g.slot_ref for g in forward.groups), head_workspace_bytes=self._plan_head_workspace_bytes(forward), head_backward_traced=self._plan_head_backward_traced(forward), @@ -1698,6 +2009,7 @@ def _estimate_required_memory_bytes_from_values( gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), group_routed_rows: tuple[int, ...] | None = None, + group_layouts: tuple[_GroupLayout, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, head_backward_traced: bool = False, @@ -1786,7 +2098,11 @@ def _estimate_required_memory_bytes_from_values( ) retained, workspace = ( self._checkpoint_memory_floor( - group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows + group_rows, + slot_refs, + gdn_segments, + routed_rows=group_routed_rows, + layouts=group_layouts, ) if checkpoint_memory is None else checkpoint_memory @@ -1797,11 +2113,15 @@ def _estimate_required_memory_bytes_from_values( # The input gradient, the backward's other end and cold transients, # staged as _subforward_cost. gradient = self._checkpoint_input_gradient_bytes(group_rows, slot_refs) - adapter_gradient = self._checkpoint_adapter_gradient_bytes( - self._checkpoint_gradient_groups(group_rows, slot_refs) + adapter_gradient = self._checkpoint_adapter_gradient_extra( + (retained, workspace), + group_rows, + slot_refs, + group_routed_rows, + group_layouts, ) head_stage = self._checkpoint_head_stage_bytes( - head_workspace_bytes, gradient, group_rows, slot_refs + head_workspace_bytes, gradient, group_rows, slot_refs, group_layouts ) peak += adapter_gradient if head_stage is not None: diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 5e2b7bfa8..aa98f07d2 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -59,6 +59,7 @@ _FlatForwardPlan, _ForwardGroupPlan, _ForwardRefusal, + _GroupLayout, _LayoutKey, _MemoryCheck, _MemorySignature, @@ -433,6 +434,7 @@ def _split_chunk_lower_cost( head_workspace_bytes = 0 head_traced: list[bool | None] = [] group_rows: list[tuple[int, bool]] = [] + group_physical_rows: list[int] = [] for (_slot, grad_enabled), group_indices in groups: estimated = _impl.estimate_prefix_tree_packed_tokens( (rows[index] for index in group_indices), @@ -440,6 +442,7 @@ def _split_chunk_lower_cost( ) assert estimated is not None # rows are CPU copies physical_rows = self._physical_tokens(estimated) + group_physical_rows.append(physical_rows) packed_tokens += physical_rows # The most loaded CP rank holds at least an even share. cp = max(1, self._topology_key()[2]) @@ -474,6 +477,16 @@ def _split_chunk_lower_cost( slot_groups=tuple(key for key, _ in groups), ) logical_tokens = _impl._active_logical_tokens(requests) + # Where exact plans price each rank's layouts, bound them from below + # the same way rather than with the busiest-rank widths. + layouts = ( + self._minimum_layouts(group_physical_rows, signature.topology[2]) + if self._layout_pricing_supported( + signature.topology, + gradient_groups=all(grad for _, grad in group_rows), + ) + else None + ) # Whether the exact plan stages its head can depend on projected rows # these bounds leave open; the lower of both prices bounds either. cost = min( @@ -484,6 +497,7 @@ def _split_chunk_lower_cost( signature=signature, logical_tokens=logical_tokens, group_rows=tuple(group_rows), + group_layouts=layouts, slot_refs=tuple(ref for (ref, _), _ in groups), head_workspace_bytes=head_workspace_bytes, head_backward_traced=traced, @@ -625,6 +639,85 @@ def _plan_group_routed_rows( ) +def _plan_group_layouts( + self: TrainerRank, plan: _FlatForwardPlan +) -> tuple[_GroupLayout, ...] | None: + """Every rank's CP layouts per group, where layout pricing is modeled. + + Only CP2 at TP1/PP1 with gradient groups, ART's CP core attention with + no softmax offset, and GDN layers marked with island boundaries; the + executor's retained set is validated there. Elsewhere ``None`` keeps + the busiest-rank pricing. + """ + if not plan.groups or not self._layout_pricing_supported( + plan.signature.topology, + gradient_groups=all(group.grad_enabled for group in plan.groups), + ): + return None + started = _impl.time.perf_counter() + try: + return self._compute_group_layouts(plan) + finally: + # Planning work: every rank's CP plan, cached by planning key. + self._planning_seconds_accum += _impl.time.perf_counter() - started + + +def _compute_group_layouts( + self: TrainerRank, plan: _FlatForwardPlan +) -> tuple[_GroupLayout, ...]: + geometry = self._geometry + from art.megatron.context_parallel.executor import retained_stage_record_bytes + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.flex_attn.compiled import flash_sparse_block_size_for_head_dim + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + _gdn_planner_config_for_provider, + ) + + topology = self._topology() + handler = self.runtime.model_support_handler + config = _context_parallel_config_for_provider( + self.runtime.provider, self.device, handler + ) + head = int(geometry.kv_channels) + block = flash_sparse_block_size_for_head_dim( + head_dim=head, head_dim_v=head, device=self.device + ) + layouts = [] + for group in plan.groups: + batch = _impl._pad_packed_batch(group.packed, multiple=int(topology.tp)) + attention, gdn, rank_plans = context_parallel_rank_layouts( + group_ids=batch.group_ids, + parent_ids=batch.parent_ids, + topology=topology, + config=config, + original_seq_len=int(batch.tokens.shape[1]), + build_gdn_execution_spec=handler.build_gdn_execution_spec, + gdn_planner_config=_gdn_planner_config_for_provider( + self.runtime.provider, handler + ), + ) + layouts.append( + _impl._GroupLayout( + attention_rows=attention, + gdn_rows=gdn, + attention_retained=tuple( + retained_stage_record_bytes( + rank_plan, + q_heads=int(geometry.num_attention_heads), + kv_heads=int(geometry.num_query_groups), + head_dim=head, + value_head_dim=head, + element_size=self._param_dtype_size, + block_size=block, + ) + for rank_plan in rank_plans + ), + ) + ) + return tuple(layouts) + + def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -634,6 +727,7 @@ def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: gdn_segments=plan.grad_segment_count, group_rows=self._plan_group_rows(plan), group_routed_rows=self._plan_group_routed_rows(plan), + group_layouts=self._plan_group_layouts(plan), slot_refs=tuple(g.slot_ref for g in plan.groups), head_workspace_bytes=self._plan_head_workspace_bytes(plan), head_backward_traced=self._plan_head_backward_traced(plan), diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index b1df930f1..6b02c7274 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -695,6 +695,7 @@ def replay( name in arguments for name in ( "slot_refs", + "group_layouts", "head_workspace_bytes", "checkpoint_floor", ) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 68890c1c3..a38519916 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -128,6 +128,16 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_recomputed_mixer_bytes_per_token", "_mixer_activation_widths", "_triton_min_rows", + "_plan_group_layouts", + "_compute_group_layouts", + "_layout_pricing_supported", + "_layer_gdn_inputs", + "_layout_checkpoint_floor", + "_layout_checkpoint_rank_floors", + "_layout_layer_boundaries", + "_layout_adapter_gradient_bytes", + "_checkpoint_adapter_gradient_extra", + "_recomputed_mixer_widths", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) @@ -178,8 +188,10 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: return [coefficient, [list(stage) for stage in stages], shared] has_grad = any(g.grad_enabled for g in plan.groups) - for group, (physical_rows, _) in zip( - plan.groups, rank._plan_group_rows(plan), strict=True + # Every rank's CP layouts come from the live runtime's CP configuration. + plan_layouts = rank._plan_group_layouts(plan) + for index, (group, (physical_rows, _)) in enumerate( + zip(plan.groups, rank._plan_group_rows(plan), strict=True) ): segments = group.packed.segments remaining -= len(segments) @@ -258,6 +270,17 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: raise ValueError("runtime_shape_inventory_over_limit") reserve(128 + 12 * len(name) + 24 * len(pending)) adapter = {"kind": kind, "name": name, "pending": [int(v) for v in pending]} + layout = None + if plan_layouts is not None: + chosen = plan_layouts[index] + if len(chosen.attention_rows) > 64: + raise ValueError("runtime_shape_inventory_over_limit") + reserve(96 + 72 * len(chosen.attention_rows)) + layout = { + "attention_rows": list(chosen.attention_rows), + "gdn_rows": None if chosen.gdn_rows is None else list(chosen.gdn_rows), + "attention_retained": list(chosen.attention_retained), + } model = _gdn_memory.model_shapes(rank, group.slot_ref) if has_grad else None if model is not None: if len(model[1]) > 1024: @@ -278,6 +301,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "adapter": adapter, "moe_covered": group.grad_enabled and rank._moe_recompute_covered_for(group.slot_ref), + "layout": layout, "gdn": None if model is None else { @@ -300,8 +324,13 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: } ) layers = _memory._checkpoint_layers(rank, rank._plan_group_rows(plan)) + # Per-rank layout pricing reads each decoder layer's input layout. + inputs = rank._layer_gdn_inputs() if plan_layouts is not None and layers else () + if len(inputs) > MAX_LAYERS: + raise ValueError("runtime_shape_inventory_over_limit") + reserve(8 * len(inputs)) facts = { - "version": 3, + "version": 4, "checkpoint_layers": layers, "checkpoint_moe_bytes_per_token": rank._checkpoint_moe_bytes_per_token(), # Model and process readers of the recomputed layer and staged head. @@ -312,6 +341,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: # Only a checkpointed Megatron decoder stages the head and has this. "backward_row_state_bytes": rank._backward_row_state_bytes() if layers else 0, "triton_min_rows": rank._triton_min_rows(), + "layer_gdn_inputs": list(inputs), "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -343,12 +373,13 @@ def integer(value: Any, *, minimum: int = 0) -> None: "te_workspace_growth_bytes", "backward_row_state_bytes", "triton_min_rows", + "layer_gdn_inputs", "head_vocabulary", "head_target_backward", "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 3: + if type(facts["version"]) is not int or facts["version"] != 4: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", @@ -383,6 +414,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "head_target_rows", "adapter", "moe_covered", + "layout", "gdn", }, ) @@ -436,6 +468,24 @@ def integer(value: Any, *, minimum: int = 0) -> None: integer(terms[2]) if terms[2] > terms[0]: raise ValueError("invalid MoE terms") + layout = group["layout"] + if layout is not None: + fields(layout, {"attention_rows", "gdn_rows", "attention_retained"}) + ranks = layout["attention_rows"] + gdn_rows = layout["gdn_rows"] + if ( + not group["grad"] + or type(ranks) is not list + or not 0 < len(ranks) <= 64 + or (gdn_rows is not None and type(gdn_rows) is not list) + or len(gdn_rows or ranks) != len(ranks) + or type(layout["attention_retained"]) is not list + or len(layout["attention_retained"]) != len(ranks) + ): + raise ValueError("invalid CP layout facts") + reserve(96 + 72 * len(ranks)) + for value in (*ranks, *(gdn_rows or ()), *layout["attention_retained"]): + integer(value) adapter = group["adapter"] if adapter is not None: fields(adapter, {"kind", "name", "pending"}) @@ -490,6 +540,23 @@ def integer(value: Any, *, minimum: int = 0) -> None: } if len(kinds) > 1: raise ValueError("invalid adapter gradient facts") + # Live layouts cover every group of a plan on the same ranks, or none. + layouts = [group["layout"] for group in groups] + laid_out = any(layout is not None for layout in layouts) + if laid_out and ( + any(layout is None for layout in layouts) + or len({len(layout["attention_rows"]) for layout in layouts}) != 1 + ): + raise ValueError("invalid CP layout facts") + inputs = facts["layer_gdn_inputs"] + if ( + type(inputs) is not list + or len(inputs) > MAX_LAYERS + or len(inputs) != (facts["checkpoint_layers"] if laid_out else 0) + or any(type(value) is not bool for value in inputs) + ): + raise ValueError("invalid layer input layout facts") + reserve(8 * len(inputs)) if len(json.dumps(facts, separators=(",", ":"))) > _MAX_BYTES: raise ValueError("runtime_facts_over_limit") @@ -664,10 +731,11 @@ def _checkpoint_memory_floor( slot_refs: Any = None, gdn_segments: int = 0, routed_rows: Any = None, + layouts: Any = None, ) -> tuple[int, int]: if self._facts is None: return _memory._checkpoint_memory_floor( - self, group_rows, slot_refs, gdn_segments, routed_rows + self, group_rows, slot_refs, gdn_segments, routed_rows, layouts ) return _memory._checkpoint_floor_from_facts( self, @@ -676,8 +744,14 @@ def _checkpoint_memory_floor( gdn_segments, self._facts["checkpoint_layers"], routed_rows, + layouts, ) + def _layer_gdn_inputs(self) -> tuple[bool, ...]: + if self._facts is None: + return _memory._layer_gdn_inputs(self) + return tuple(self._facts["layer_gdn_inputs"]) + def _moe_recompute_covered_for(self, slot_ref: Any) -> bool: if self._facts is None: return _memory._moe_recompute_covered_for(self, slot_ref) @@ -723,6 +797,20 @@ def runtime_arguments( raise ValueError("runtime facts disagree with selected routed rows") if type(arguments.get("head_backward_traced")) is not bool: raise ValueError("head backward staging is not recorded") + layouts = None + if groups[0]["layout"] is not None: + if len(groups[0]["layout"]["attention_rows"]) != self._topology_key()[2]: + raise ValueError("CP layout facts disagree with the recorded topology") + layouts = tuple( + _impl._GroupLayout( + attention_rows=tuple(g["layout"]["attention_rows"]), + gdn_rows=None + if g["layout"]["gdn_rows"] is None + else tuple(g["layout"]["gdn_rows"]), + attention_retained=tuple(g["layout"]["attention_retained"]), + ) + for g in groups + ) head = max( max( _memory._dense_head_bytes(facts["head_vocabulary"], g["head_rows"]), @@ -755,6 +843,7 @@ def runtime_arguments( **arguments, "group_rows": rows, "group_routed_rows": routed, + "group_layouts": layouts, "slot_refs": tuple(range(len(groups))), "head_workspace_bytes": head, "checkpoint_floor": (retained, workspace), diff --git a/tests/unit/test_context_parallel_retained_bytes.py b/tests/unit/test_context_parallel_retained_bytes.py new file mode 100644 index 000000000..18f0f4b48 --- /dev/null +++ b/tests/unit/test_context_parallel_retained_bytes.py @@ -0,0 +1,157 @@ +"""Size-only retained-set mirror of the CP attention executor's recompute path.""" + +import pytest + +pytest.importorskip("triton") + +from art.megatron.context_parallel.executor import ( # noqa: E402 + minimum_retained_bytes_per_row, + retained_stage_record_bytes, +) +from art.megatron.context_parallel.types import ( # noqa: E402 + RankRuntimePlan, + StagePlan, + TokenRange, +) + + +# Qwen3.6-35B-A3B attention: 16 query heads, 2 KV heads of 256, BF16, and the +# H200 flash block for a 256-wide head (128 query, 64 key rows). +def retained(runtime_plan, *, q_heads=16, kv_heads=2): + return retained_stage_record_bytes( + runtime_plan, + q_heads=q_heads, + kv_heads=kv_heads, + head_dim=256, + value_head_dim=256, + element_size=2, + block_size=(128, 64), + ) + + +def minimum(*, q_heads=16, kv_heads=2): + return minimum_retained_bytes_per_row( + q_heads=q_heads, + kv_heads=kv_heads, + head_dim=256, + value_head_dim=256, + element_size=2, + ) + + +# Flex's output keeps its own LSE beside the normalized copy it returns; +# logical-length copies keep one. +Q, KV, FLEX, OUT, TAPE = 8192, 2048, 8192 + 2 * 64, 8192 + 64, 16 * 257 * 4 + + +def stage(index, *, local, q, k, q_len=None, k_len=None, own=None, source=0): + q_ranges = ( + (TokenRange(0, q),) if own is None or q == own else (TokenRange(1, 1 + q),) + ) + return StagePlan( + stage_index=index, + source_rank=source, + is_local_stage=local, + slices=("slice",), # ty: ignore[invalid-argument-type] + owner_local_q_ranges=q_ranges if q else (), + owner_local_k_ranges=(TokenRange(0, k),) if k else (), + q_len=q if q_len is None else q_len, + k_len=k if k_len is None else k_len, + ) + + +def plan(own, *stages): + return RankRuntimePlan( + rank=0, + original_seq_len=own, + token_layout_index=None, # ty: ignore[invalid-argument-type] + local_valid_lengths=(own,), + local_row_ranges=(TokenRange(0, own),), + stage_plans=stages, + remote_dkv_reduce_plan=None, # ty: ignore[invalid-argument-type] + ) + + +def test_aligned_single_stage_copies_views_without_output_copies(): + # Real-data rank 0: one aligned local stage. Contiguous copies of the + # permuted Q/K/V views, flex output and LSE; no padding, copies or tape. + rows = 52480 + kept = retained(plan(rows, stage(0, local=True, q=rows, k=rows))) + assert kept == Q * rows + KV * rows + FLEX * rows # 0.971 GB traced + assert kept == rows * minimum() + + +def test_unaligned_single_stage_pads_and_copies_the_logical_output(): + # The tail chunk on the single-stage rank: the planner rounds both lengths + # up to its 128-row block, so both pad. + rows = 52481 + stage_len = 411 * 128 + kept = retained( + plan( + rows, stage(0, local=True, q=rows, k=rows, q_len=stage_len, k_len=stage_len) + ) + ) + assert kept == Q * stage_len + KV * stage_len + FLEX * stage_len + OUT * rows + + +def test_tiny_stage_pads_to_two_blocks(): + kept = retained(plan(5, stage(0, local=True, q=5, k=5))) + assert kept == Q * 256 + KV * 128 + FLEX * 256 + OUT * 5 + + +def test_single_head_views_may_still_be_copied(): + # One head does not make a view of a fused QKV split contiguous, so the + # mirror still charges the copies; only the lower bound leaves them out. + rows = 1024 + kept = retained( + plan(rows, stage(0, local=True, q=rows, k=rows)), q_heads=1, kv_heads=1 + ) + assert kept == (512 + 1024 + 512 + 2 * 4) * rows + assert minimum(q_heads=1, kv_heads=1) == 512 + 2 * 4 + + +def test_full_query_remote_stage_keeps_fetch_buffers_and_a_merge_tape(): + # Real-data rank 1: a local stage and a remote stage over all of its + # queries. Planner lengths round up, so both stages pad. + own, remote_k = 44314, 16504 + local = stage(0, local=True, q=own, k=own, q_len=44352, k_len=44352) + remote = stage(1, local=False, q=own, k=remote_k, q_len=44352, k_len=16512) + kept = retained(plan(own, local, remote)) + q_pad = 347 * 128 # 44,416, as flex's traced output size shows + local_bytes = Q * q_pad + KV * 44352 + FLEX * q_pad + OUT * own + remote_bytes = ( + Q * q_pad + KV * remote_k + KV * 16512 + FLEX * q_pad + OUT * own + TAPE * own + ) + assert kept == local_bytes + remote_bytes + assert 3.05e9 < kept < 3.10e9 # 3.07 GB traced + + +def test_partial_query_remote_stage_keeps_its_gather_and_a_partial_tape(): + # Random-data rank 1: a large local stage and a 768-row remote stage. + own = 105153 + local = stage(0, local=True, q=own, k=own, q_len=105216, k_len=105216) + remote = stage(1, local=False, q=768, k=20608, own=own) + kept = retained(plan(own, local, remote)) + local_bytes = Q * 105216 + KV * 105216 + FLEX * 105216 + OUT * own + # Aligned: the query gather and fetch buffers feed flex without copies; the + # partial tape also keeps its int64 row index. + remote_bytes = Q * 768 + KV * 20608 + FLEX * 768 + (TAPE + 8) * 768 + assert kept == local_bytes + remote_bytes + + +def test_empty_remote_stage_and_missing_local_stage(): + rows = 1024 + empty = stage(1, local=False, q=0, k=0) + alone = retained(plan(rows, stage(0, local=True, q=rows, k=rows), empty)) + assert alone == Q * rows + KV * rows + FLEX * rows + # Without a local stage, the first ready remote stage records no tape; the + # mirror drops the smallest so it never under-counts the order. + small = stage(1, local=False, q=256, k=256, own=rows) + full = stage(2, local=False, q=rows, k=512) + both = retained(plan(rows, small, full)) + without_small_tape = retained(plan(rows, full)) + # The dropped tape's int64 index stays: the executor keeps it regardless. + assert both - without_small_tape == ( + Q * 256 + KV * 256 + FLEX * 256 + 8 * 256 + TAPE * rows + ) + assert retained(plan(rows)) == 0 diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 23a93a29a..78316dde6 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -571,3 +571,105 @@ def test_grouped_replay_refuses_unconsumed_metadata(field, monkeypatch, tmp_path report["replay"][field].append(deepcopy(report["replay"][field][0])) with pytest.raises(ValueError, match="unused"): reports.replay(report) + + +def layout_report(monkeypatch, tmp_path): + """A grouped CP2 report priced on every rank's own layouts.""" + from test_trainer_rank_checkpoint_memory import rank as checkpoint_rank + from test_trainer_rank_layout_memory import _requests, art_cp + + rank = art_cp(checkpoint_rank(), monkeypatch) + del rank._topology_key # the stock reader, over the fixture's CP2 topology + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward(_requests()) + assert rank._plan_group_layouts(plan) is not None + report, costs = emitted(rank, plan, tmp_path) + return report, costs, rank + + +def test_cp_layouts_are_replayed_and_frozen(monkeypatch, tmp_path): + original, costs, rank = layout_report(monkeypatch, tmp_path) + facts = original["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + (group,) = facts["groups"] + assert len(group["layout"]["attention_rows"]) == 2 + assert len(facts["layer_gdn_inputs"]) == facts["checkpoint_layers"] == 40 + assert any(facts["layer_gdn_inputs"]) and not all(facts["layer_gdn_inputs"]) + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + assert actual["estimates"][0]["required_bytes"] == costs[0].required + # The live model's layers cannot change the replayed answer. + for layer in rank.runtime.model[0].decoder.layers: + layer._art_gdn_island_boundary = None + assert reports.replay(original) == actual + + def changed(edit): + report = deepcopy(original) + edit(report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]) + return reports.replay(report)["estimates"][0] + + # Each frozen layout fact reaches the per-rank pricing. + for edit in ( + lambda f: f["groups"][0]["layout"]["attention_rows"].__setitem__(1, 10**5), + lambda f: f["groups"][0]["layout"]["attention_retained"].__setitem__(0, 10**12), + lambda f: f.__setitem__("layer_gdn_inputs", [True] * 40), + ): + result = changed(edit) + assert not result["matches"] + assert result["required_bytes"] != actual["estimates"][0]["required_bytes"] + + +@pytest.mark.parametrize( + "change", + [ + "no_grad", + "gdn_length", + "value", + "inputs_length", + "inputs_type", + "ranks", + "argument", + ], +) +def test_cp_layout_fact_validation_rejects_forged_input(change, monkeypatch, tmp_path): + report, _, _ = layout_report(monkeypatch, tmp_path) + item = report["replay"]["memory_replay"]["estimates"][0] + facts = item["runtime_facts"] + layout = facts["groups"][0]["layout"] + if change == "no_grad": + facts["groups"][0]["grad"] = False + facts["groups"][0]["moe_covered"] = False + facts["groups"][0]["adapter"] = None + message = "invalid CP layout facts" + elif change == "gdn_length": + layout["gdn_rows"].append(1) + message = "invalid CP layout facts" + elif change == "value": + layout["attention_retained"][0] = 1.5 + message = "invalid runtime dimension" + elif change == "inputs_length": + facts["layer_gdn_inputs"].pop() + message = "invalid layer input layout facts" + elif change == "inputs_type": + facts["layer_gdn_inputs"][0] = 1 + message = "invalid layer input layout facts" + elif change == "ranks": + for key in ("attention_rows", "gdn_rows", "attention_retained"): + layout[key].append(1) + message = "recorded topology" + else: + item["arguments"]["group_layouts"] = [dict(layout)] + message = "immutable runtime" + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +@pytest.mark.parametrize("name", ["_plan_group_layouts", "_layer_gdn_inputs"]) +def test_custom_layout_reader_is_explicitly_incomplete(name, monkeypatch): + from art.trainer_rank import _planner_replay + + rank = _rank(monkeypatch) + plan = rank._plan_flat_forward([_request(1)]) + original = getattr(rank, name) + monkeypatch.setattr(rank, name, lambda *args: original(*args)) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index 3481bfd6d..66b4bf812 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -33,6 +33,7 @@ def assert_plan_values(rank, plan, values): assert values["gdn_segments"] == plan.grad_segment_count assert values["group_rows"] == rank._plan_group_rows(plan) assert values["group_routed_rows"] == rank._plan_group_routed_rows(plan) + assert values["group_layouts"] == rank._plan_group_layouts(plan) assert values["head_workspace_bytes"] == rank._plan_head_workspace_bytes(plan) assert values["checkpoint_floor"] == _gdn_memory.plan_floor(rank, plan) assert values["retained_tokens"] == rank._plan_retained_tokens(plan) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index 78fcb7f73..223c2dd54 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -534,10 +534,11 @@ def test_routed_rows_move_only_the_routed_moe_part(): fewer = r._checkpoint_memory_floor(((local, True),), None, routed_rows=(routed,)) assert fewer[0] == retained assert workspace - fewer[1] == (local - routed) * (188416 - 8192) - # Never more routed rows than local ones. - assert r._moe_workspace_bytes( - 10, routed_rows=20, checkpoint_grad=True - ) == r._moe_workspace_bytes(10, checkpoint_grad=True) + # A rank can receive more routed rows than it holds: all are priced, and + # the shared part stays on the routed count. + assert r._moe_workspace_bytes(10, routed_rows=20, checkpoint_grad=True) == ( + 20 * 188416 + ) r._moe_gradient_shared_bytes = 188417 with pytest.raises(ValueError, match="shared-expert"): r._moe_workspace_bytes(10, checkpoint_grad=True) diff --git a/tests/unit/test_trainer_rank_layout_memory.py b/tests/unit/test_trainer_rank_layout_memory.py new file mode 100644 index 000000000..a8394276d --- /dev/null +++ b/tests/unit/test_trainer_rank_layout_memory.py @@ -0,0 +1,595 @@ +"""CP2 layout-aware recompute pricing: per-rank ledgers and gates, CPU only.""" + +from dataclasses import replace +from types import SimpleNamespace + +import pytest +from test_trainer_rank_checkpoint_memory import rank +import torch + +from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import _TE_CUBLAS_WORKSPACE_BYTES, Unset, _GroupLayout + +H = 2048 * 2 + + +def qwen36(r): + """Qwen3.6-35B-A3B's 3:1 GDN/attention pattern and mixer geometry at CP2.""" + r._topology_key = lambda: (1, 1, 2, 1) + r._geometry = replace( + r._geometry, + num_attention_heads=16, + num_query_groups=2, + kv_channels=256, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=32, + gdn_value_head_dim=128, + ) + r._attention_output_gate = True + r._gdn_layers = 30 + for index, layer in enumerate(r.runtime.model[0].decoder.layers): + is_gdn = index % 4 != 3 + after_gdn = index > 0 and (index - 1) % 4 != 3 + layer._art_gdn_island_boundary = SimpleNamespace( + is_gdn=is_gdn, input_layout="gdn" if is_gdn and after_gdn else "attention" + ) + return r + + +def test_boundaries_follow_each_layer_inputs_layout(): + # Traced real-data CP2 ranks: 20 inputs in each layout, and boundaries of + # 8.248 and 7.612 GB. The busiest attention rank is not the busiest overall. + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + layout = _GroupLayout((52480, 44314), (48194, 48600), (0, 0)) + retained, _ = r._layout_checkpoint_floor((None,), (48397,), (layout,)) + assert retained == H * (20 * 52480 + 20 * 48194) + assert H * (20 * 44314 + 20 * 48600) < retained < 40 * 52480 * H + + +def test_largest_rank_total_not_a_sum_of_rank_maxima(): + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + widths = r._recomputed_mixer_widths(stage_buffers=False) + assert widths == {"attention": 2 * (2 * 2048 + 7 * 4096 + 3 * 512), "gdn": 94208} + # Rank 1's two full-query stages keep 3.08 GB against rank 0's 0.97 GB. + layout = _GroupLayout((52480, 44314), (48194, 48600), (970_720_000, 3_079_000_000)) + retained, workspace = r._layout_checkpoint_floor((None,), (48397,), (layout,)) + + def total(rank_index): + attention = layout.attention_rows[rank_index] + assert layout.gdn_rows is not None + gdn = layout.gdn_rows[rank_index] + ledger = H * (20 * attention + 20 * gdn) + stages = ( + attention * (widths["attention"] + 2 * H) + + layout.attention_retained[rank_index] + + r._moe_workspace_bytes( + attention, routed_rows=48397, checkpoint_grad=True + ), + gdn * (widths["gdn"] + 2 * H) + + r._moe_workspace_bytes(gdn, routed_rows=48397, checkpoint_grad=True), + ) + return ledger + max(stages) + + assert retained + workspace == max(total(0), total(1)) + _TE_CUBLAS_WORKSPACE_BYTES + # Rank 1 has fewer rows but is the largest total: its attention runs two + # full-query stages. Pricing it on its own rows is about neutral against + # the busiest-rank floor, which over-counts boundaries and GDN instead. + assert total(1) > total(0) + busiest = sum( + r._checkpoint_memory_floor(((52480, True),), None, routed_rows=(48397,)) + ) + assert abs(retained + workspace - busiest) < 0.01 * busiest + + +def test_routed_rows_above_a_ranks_own_rows_are_all_priced(): + r = qwen36(rank()) + layers = r.runtime.model[0].decoder.layers + few = _GroupLayout((100, 100), (100, 100), (0, 0)) + many = _GroupLayout((100, 100), (100, 100), (0, 0)) + low = sum(r._layout_checkpoint_floor((None,), (100,), (few,))) + high = sum(r._layout_checkpoint_floor((None,), (400,), (many,))) + assert high - low == 300 * 188416 + + +def test_no_grad_groups_keep_busiest_rank_pricing(): + r = qwen36(rank()) + layout = _GroupLayout((10, 8), (9, 9), (0, 0)) + groups = ((10, True), (12, False)) + assert r._checkpoint_memory_floor( + groups, None, routed_rows=(9, 12), layouts=(layout, layout) + ) == r._checkpoint_memory_floor(groups, None, routed_rows=(9, 12)) + + +def _plan(r, *, no_grad=False): + plan = r._plan_flat_forward( + [ + ForwardInput( + input_tokens=torch.arange(64), hidden_states=True, no_grad=no_grad + ) + ] + ) + return replace(plan, signature=replace(plan.signature, topology=(1, 1, 2, 1))) + + +@pytest.mark.parametrize( + "topology", [(1, 1, 1, 1), (1, 2, 2, 1), (1, 1, 4, 1), (1, 1, 2, 2)] +) +def test_layout_pricing_is_cp2_tp1_pp1_only(topology): + r = qwen36(rank()) + plan = _plan(r) + plan = replace(plan, signature=replace(plan.signature, topology=topology)) + assert r._plan_group_layouts(plan) is None + + +def test_layout_pricing_needs_gradient_groups_and_art_cp_attention(): + r = qwen36(rank()) + assert r._plan_group_layouts(_plan(r, no_grad=True)) is None + # The stub decoder's attention layers carry no ART CP core attention. + assert r._plan_group_layouts(_plan(r)) is None + # A GDN model without island boundaries is not modeled either. + for layer in r.runtime.model[0].decoder.layers: + del layer._art_gdn_island_boundary + assert r._plan_group_layouts(_plan(r)) is None + + +def test_plan_cost_and_admission_use_the_same_layouts(monkeypatch): + r = qwen36(rank()) + plan = _plan(r) + layout = _GroupLayout((40, 24), (32, 32), (1_000_000, 3_000_000)) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((40, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (32,)) + monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: 32) + cost = r._plan_cost(plan) + # Admission adds CP output coexistence the plan cost leaves out (as it did + # before layouts); the layouts themselves enter both the same way. + gap = r._memory_check(plan).estimated_required_bytes - cost.required + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: None) + assert gap == r._memory_check(plan).estimated_required_bytes - ( + r._plan_cost(plan).required + ) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + retained, _ = r._layout_checkpoint_floor( + (None,), + r._plan_group_routed_rows(plan), + (layout,), + ) + assert cost.checkpoint_retained == plan.output_bytes + retained + + +def art_cp(r, monkeypatch): + """Qwen3.6 at CP2 with ART's CP core attention and a CPU planning config.""" + from art.megatron.context_parallel.core_attention import ( + ArtContextParallelCoreAttention, + ) + from art.megatron.context_parallel.types import ParallelTopology + + r = qwen36(r) + for index, layer in enumerate(r.runtime.model[0].decoder.layers): + if index % 4 == 3: + core = ArtContextParallelCoreAttention.__new__( + ArtContextParallelCoreAttention + ) + torch.nn.Module.__init__(core) + core.softmax_offset = None + layer.self_attention = torch.nn.Module() + layer.self_attention.core_attention = core + provider = r.runtime.provider + provider.kv_channels = 256 + provider.ffn_hidden_size = 8192 + provider.linear_num_key_heads = 16 + provider.linear_num_value_heads = 32 + provider.linear_key_head_dim = 128 + provider.linear_value_head_dim = 128 + provider.params_dtype = torch.bfloat16 + r.runtime.model_support_handler = SimpleNamespace( + build_gdn_execution_spec=True, + context_parallel_workload_profile=lambda provider: None, + ) + monkeypatch.setattr(r, "_topology", lambda: ParallelTopology(tp=1, cp=2)) + return r + + +def _requests(lengths=(900, 700, 500)): + start, requests = 0, [] + for n in lengths: + requests.append( + ForwardInput( + input_tokens=torch.arange(start, start + n), hidden_states=True + ) + ) + start += n + return requests + + +def test_rank_layouts_are_the_executors_plans(monkeypatch): + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.context_parallel.types import ( + ContextParallelConfig, + ParallelTopology, + ) + from art.megatron.prefix_tree_packing import prefix_tree_pack + + packed = prefix_tree_pack([r.input_tokens for r in _requests()], max_depth=1) + attention, gdn, plans = context_parallel_rank_layouts( + group_ids=packed.group_ids, + parent_ids=packed.parent_ids, + topology=ParallelTopology(tp=1, cp=2), + config=ContextParallelConfig(), + original_seq_len=int(packed.tokens.shape[1]), + build_gdn_execution_spec=True, + ) + total = int(packed.tokens.numel()) + assert sum(attention) == total and gdn is not None and sum(gdn) == total + # The ledger's attention rows and the mirror's own rows are the same count. + assert attention == tuple(plan.local_valid_lengths[0] for plan in plans) + + +def test_gated_plans_price_every_rank_from_its_stage_plan(monkeypatch): + from art.megatron.context_parallel.executor import retained_stage_record_bytes + from art.megatron.context_parallel.runtime import context_parallel_rank_layouts + from art.megatron.context_parallel.types import ParallelTopology + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + ) + + r = art_cp(rank(), monkeypatch) + plan = _plan_with(r, _requests()) + (layout,) = r._plan_group_layouts(plan) + assert sum(layout.attention_rows) == plan.packed_tokens + assert layout.gdn_rows is not None and sum(layout.gdn_rows) == plan.packed_tokens + # Each rank's retention is the executor mirror over that rank's own plan. + (group,) = plan.groups + _, _, rank_plans = context_parallel_rank_layouts( + group_ids=group.packed.group_ids, + parent_ids=group.packed.parent_ids, + topology=ParallelTopology(tp=1, cp=2), + config=_context_parallel_config_for_provider( + r.runtime.provider, r.device, r.runtime.model_support_handler + ), + original_seq_len=int(group.packed.tokens.shape[1]), + build_gdn_execution_spec=True, + ) + assert layout.attention_retained == tuple( + retained_stage_record_bytes( + rank_plan, + q_heads=16, + kv_heads=2, + head_dim=256, + value_head_dim=256, + element_size=2, + block_size=(128, 128), # the CPU device's flex block + ) + for rank_plan in rank_plans + ) + assert all(retained > 0 for retained in layout.attention_retained) + + +def test_softmax_offset_leaves_the_busiest_rank_floor(monkeypatch): + r = art_cp(rank(), monkeypatch) + plan = _plan_with(r, _requests()) + assert r._plan_group_layouts(plan) is not None + for layer in r.runtime.model[0].decoder.layers: + core = getattr(getattr(layer, "self_attention", None), "core_attention", None) + if core is not None: + core.softmax_offset = torch.zeros(16) + assert r._plan_group_layouts(plan) is None + + +@pytest.mark.parametrize("per_layer", [0, 6000 * H]) +@pytest.mark.parametrize( + "lengths", + [(2048, 1536, 1024, 512), (4099, 3, 5, 7), (1, 2, 3, 4, 5, 6, 7), (8191,)], +) +def test_split_lower_bound_stays_below_the_layout_cost(monkeypatch, lengths, per_layer): + # Even and skewed CP splits, odd row counts, and a single long sequence; + # with pending policy gradients, even shares still price below each + # rank's paired extra. + r = art_cp(rank(), monkeypatch) + if per_layer: + _with_policy_gradients(monkeypatch, r, (per_layer,) * 40 + (0,)) + requests = _requests(lengths) + plan = _plan_with(r, requests) + lower = r._split_chunk_lower_cost( + requests, tuple(item.input_tokens for item in requests), checkpoint=Unset + ) + assert lower.required <= r._plan_cost(plan).required + # It prices even-share layouts at the least attention state. + assert r._layout_pricing_supported((1, 1, 2, 1), gradient_groups=True) + + +@pytest.mark.parametrize( + "lengths", [(2048, 1536, 1024, 512), (4099, 3, 5, 7), (8191, 64)] +) +def test_split_lower_bound_stays_below_two_slots_layout_cost(monkeypatch, lengths): + # Two gradient groups of different slots (the first request alone, then + # the rest), each with pending gradients. + from art.megatron.lora import LoRASlotRef + + r = art_cp(rank(), monkeypatch) + base = LoRASlotRef("checkpoint", None) + monkeypatch.setattr( + r, + "_group_active_request_indices", + lambda requests, **_: ( + ((None, True), (0,)), + ((base, True), tuple(range(1, len(requests)))), + ), + ) + slots = (LoRASlotRef("checkpoint", "policy"), LoRASlotRef("checkpoint", "other")) + groups = r._checkpoint_gradient_groups + monkeypatch.setattr( + r, + "_checkpoint_gradient_groups", + lambda group_rows, slot_refs: tuple( + (slot, boundaries) + for slot, (_, boundaries) in zip(slots, groups(group_rows, slot_refs)) + ), + ) + pending = { + (slots[0],): (6000 * H,) * 40 + (0,), + (slots[1],): (300 * H,) * 40 + (7 * H,), + } + monkeypatch.setattr( + r, "_pending_adapter_gradient_bytes", lambda refs: pending.get(tuple(refs), ()) + ) + requests = _requests(lengths) + plan = _plan_with(r, requests) + assert len(plan.groups) == 2 and r._plan_group_layouts(plan) is not None + exact = r._plan_cost(plan) + assert exact.checkpoint_adapter_gradient > 0 + lower = r._split_chunk_lower_cost( + requests, tuple(item.input_tokens for item in requests), checkpoint=Unset + ) + # The even-share lower bound prices both slots' extra too, and stays below. + assert lower.checkpoint_adapter_gradient > 0 + assert lower.required <= exact.required + + +def _plan_with(r, requests): + return r._plan_flat_forward(requests) + + +def _with_policy_gradients(monkeypatch, r, pending): + """Every gradient group trains the policy slot, with ``pending`` gradients.""" + from art.megatron.lora import LoRASlotRef + + policy = LoRASlotRef("checkpoint", "policy") + groups = r._checkpoint_gradient_groups + monkeypatch.setattr( + r, + "_checkpoint_gradient_groups", + lambda group_rows, slot_refs: tuple( + (policy, boundaries) for _, boundaries in groups(group_rows, slot_refs) + ), + ) + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: pending if tuple(refs) else (), + ) + + +def _pending_adapter_rank(monkeypatch, layout, rows, per_layer=300 * H): + r = qwen36(rank()) + monkeypatch.setattr(r, "_plan_group_layouts", lambda plan: (layout,)) + monkeypatch.setattr(r, "_plan_group_rows", lambda plan: ((rows, True),)) + monkeypatch.setattr(r, "_plan_group_routed_rows", lambda plan: (rows,)) + monkeypatch.setattr(r, "_plan_hybridep_growth_bytes", lambda plan: 0) + monkeypatch.setattr(r, "_plan_retained_tokens", lambda plan: rows) + pending = (per_layer,) * 40 + (0,) + _with_policy_gradients(monkeypatch, r, pending) + layers = r.runtime.model[0].decoder.layers + gdn_inputs = [ + layer._art_gdn_island_boundary.input_layout == "gdn" for layer in layers + ] + + def boundaries(rank_index): + assert layout.gdn_rows is not None + return [ + H + * max(1, (layout.gdn_rows if is_gdn else layout.attention_rows)[rank_index]) + for is_gdn in gdn_inputs + ] + + def extra(saved): + return max( + 0, + *( + sum(pending[i:40]) + pending[40] - sum(saved[i + 1 :]) + for i in range(40) + ), + ) + + floor = r._checkpoint_memory_floor( + ((rows, True),), None, routed_rows=(rows,), layouts=(layout,) + ) + floors = r._layout_checkpoint_rank_floors((None,), (rows,), (layout,)) + return r, boundaries, extra, floor, floors + + +def test_adapter_gradients_meet_each_ranks_own_layer_boundaries(monkeypatch): + layout = _GroupLayout((400, 240), (160, 320), (0, 0)) + r, boundaries, extra, (retained, workspace), floors = _pending_adapter_rank( + monkeypatch, layout, 400 + ) + assert r._layout_layer_boundaries((layout,)) == ( + (tuple(boundaries(0)),), + (tuple(boundaries(1)),), + ) + # Every rank's boundaries sum to its own floor's. + assert [sum(boundaries(i)) for i in (0, 1)] == [floor[0] for floor in floors] + assert retained == max(floor[0] for floor in floors) + # Each rank releases its own boundaries in layer order; each extra sits on + # that rank's own floor, bounded by the floor's workspace. + expected = ( + max( + floor[0] + max(floor[1], workspace) + extra(boundaries(i)) + for i, floor in enumerate(floors) + ) + - retained + - workspace + ) + assert r._plan_cost(_plan(r)).checkpoint_adapter_gradient == expected > 0 + # An even share of the busiest rank's boundaries would misplace the peak. + assert extra([retained // 40] * 40) != expected + + +def test_a_rank_without_gdn_rows_pairs_its_extra_with_its_own_floor(monkeypatch): + # A single short sequence at CP2: every GDN row on rank 0 (Qwen3.6 2,047 + # tokens), with about 24 MB of expert LoRA gradients per layer. + layout = _GroupLayout((1536, 511), (2047, 0), (0, 0)) + r, boundaries, extra, (retained, workspace), floors = _pending_adapter_rank( + monkeypatch, layout, 2047, per_layer=6000 * H + ) + extras = [extra(boundaries(i)) for i in (0, 1)] + # Rank 1 releases almost nothing, so its extra is almost every gradient... + assert extras[1] > extras[0] and extras[1] > 6000 * H * 38 + cost = r._plan_cost(_plan(r)).checkpoint_adapter_gradient + # ...but it sits on rank 1's much smaller floor, not on rank 0's. + assert cost < extras[1] + # Every rank's own floor plus its own extra stays within the price. + for (rank_retained, rank_workspace), rank_extra in zip(floors, extras): + assert rank_retained + rank_workspace + rank_extra <= ( + retained + workspace + cost + ) + assert cost >= extras[0] + # Admission prices the same extra: its gap to the plan cost stays the CP + # output coexistence alone (up to the safety factor's integer rounding). + plan = _plan(r) + gap = r._memory_check(plan).estimated_required_bytes - r._plan_cost(plan).required + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) + assert r._plan_cost(plan).checkpoint_adapter_gradient == 0 + plain = r._memory_check(plan).estimated_required_bytes - r._plan_cost(plan).required + assert abs(gap - plain) <= 1 + + +def test_layout_gradient_groups_run_one_after_another_on_each_rank(monkeypatch): + from test_trainer_rank_adapter_gradient_memory import sequential_oracle + + from art.megatron.lora import LoRASlotRef + + r = qwen36(rank()) + # A short policy sequence beside a longer sequence of another slot. + layouts = ( + _GroupLayout((1536, 511), (2047, 0), (0, 0)), + _GroupLayout((400, 240), (160, 320), (0, 0)), + ) + group_rows = ((2047, True), (640, True)) + routed = (2047, 640) + slots = (LoRASlotRef("checkpoint", "policy"), LoRASlotRef("checkpoint", "other")) + groups = r._checkpoint_gradient_groups + monkeypatch.setattr( + r, + "_checkpoint_gradient_groups", + lambda group_rows, slot_refs: tuple( + (slot, boundaries) + for slot, (_, boundaries) in zip(slots, groups(group_rows, slot_refs)) + ), + ) + pending = { + (slots[0],): (6000 * H,) * 40 + (0,), + (slots[1],): (100 * H,) * 40 + (5 * H,), + } + walks = [] + + def pending_gradients(refs): + walks.append(tuple(refs)) + return pending.get(tuple(refs), ()) + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending_gradients) + floor = r._checkpoint_memory_floor( + group_rows, None, routed_rows=routed, layouts=layouts + ) + layers = r.runtime.model[0].decoder.layers + floors = r._layout_checkpoint_rank_floors((None, None), routed, layouts) + per_rank = r._layout_layer_boundaries(layouts) + # Every rank's groups sum to its own floor's boundaries. + assert [sum(map(sum, rank_groups)) for rank_groups in per_rank] == [ + rank_floor[0] for rank_floor in floors + ] + # On each rank, one group's backward runs after the other's, in the worst + # order; each rank's extra sits on its own floor. + extras = [ + sequential_oracle( + [(pending[(slot,)], list(b)) for slot, b in zip(slots, rank_groups)] + ) + for rank_groups in per_rank + ] + expected = max( + 0, + max( + rank_retained + max(rank_workspace, floor[1]) + extra + for (rank_retained, rank_workspace), extra in zip(floors, extras) + ) + - sum(floor), + ) + assert ( + r._checkpoint_adapter_gradient_extra(floor, group_rows, None, routed, layouts) + == expected + > 0 + ) + # Each slot's module walk runs once, not once per rank. + assert walks == [(slot,) for slot in slots] + + +def test_layout_head_stage_meets_each_ranks_other_groups(monkeypatch): + from test_trainer_rank_head_stage_memory import head_oracle + + from art.megatron.lora import LoRASlotRef + + r = qwen36(rank()) + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + # A long policy sequence whose gradients outweigh its boundaries only on + # the rank with no GDN rows, beside a short sequence of another slot. + layouts = ( + _GroupLayout((1536, 511), (2047, 0), (0, 0)), + _GroupLayout((400, 240), (160, 320), (0, 0)), + ) + group_rows = ((2047, True), (640, True)) + routed = (2047, 640) + slots = (LoRASlotRef("checkpoint", "policy"), LoRASlotRef("checkpoint", "other")) + groups = r._checkpoint_gradient_groups + monkeypatch.setattr( + r, + "_checkpoint_gradient_groups", + lambda group_rows, slot_refs: tuple( + (slot, boundaries) + for slot, (_, boundaries) in zip(slots, groups(group_rows, slot_refs)) + ), + ) + pending = { + (slots[0],): (600 * H,) * 40 + (0,), + (slots[1],): (1 * H,) * 40 + (5 * H,), + } + monkeypatch.setattr( + r, "_pending_adapter_gradient_bytes", lambda refs: pending.get(tuple(refs), ()) + ) + per_rank = r._layout_layer_boundaries(layouts) + heads = [ + head_oracle( + [(pending[(slot,)], list(b)) for slot, b in zip(slots, rank_groups)] + ) + for rank_groups in per_rank + ] + # The policy group raises the other's head on the rank that saved fewer rows. + assert heads[1] > heads[0] + head, gradient = 10**9, r._checkpoint_input_gradient_bytes(group_rows, slots) + stage = r._checkpoint_head_stage_bytes(head, gradient, group_rows, slots, layouts) + assert stage == ( + head + + 2 * gradient + + (2047 + 640) * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + + heads[1] + ) + # The busiest rank's boundaries (every layer saving each group's rows) + # would release more than the idle rank does: the per-rank walk matters. + busiest = r._checkpoint_adapter_gradient_bytes( + r._checkpoint_gradient_groups(group_rows, slots), head=True + ) + assert heads[1] > busiest diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 7bc27d8fd..136125729 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -609,7 +609,10 @@ def held(rows): monkeypatch.setattr( rank, "_checkpoint_memory_floor", - lambda rows, refs=None, segments=0, routed_rows=None: (10**7, 10**6), + lambda rows, refs=None, segments=0, routed_rows=None, layouts=None: ( + 10**7, + 10**6, + ), ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) @@ -871,6 +874,58 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( assert not torch.cuda.is_initialized() +def test_hybridep_combine_extent_floors_the_layout_path( + hybrid_checkpoint_rank, monkeypatch +): + rank = hybrid_checkpoint_rank + groups = ((2, True),) + calls = [] + + def layout_floor(refs, routed, layouts): + calls.append(layouts) + return 7, 11 + + def generic_floor(*args): + raise AssertionError("per-rank layouts must take the layout path") + + monkeypatch.setattr(rank, "_layout_checkpoint_floor", layout_floor) + # Only the busiest-rank floor prices the recomputed mixer this way. + monkeypatch.setattr(rank, "_recomputed_mixer_bytes_per_token", generic_floor) + layouts = (object(),) + te = rank._te_workspace_growth_bytes() + # Two rows round up to four; the layout floor's small workspace loses. + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + 4 * 2048 * 2 + te, + ) + marker = torch.empty(0) + rank._pending_hybridep_graphs.append(weakref.ref(marker)) + rank._hybridep_rows_high_water = 218751 + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + 218752 * 2048 * 2 + te, + ) + # The stage already carries its TE growth; the combine floor adds its own + # only when it wins, including over a stage between the two. + assert te > 0 + combine = 218752 * 2048 * 2 + for stage, workspace in ( + (combine + te - 1, combine + te), + (combine + te + 1, combine + te + 1), + ): + + def stage_layout_floor(refs, routed, layouts, stage=stage): + calls.append(layouts) + return 7, stage + + monkeypatch.setattr(rank, "_layout_checkpoint_floor", stage_layout_floor) + assert rank._checkpoint_memory_floor(groups, layouts=layouts) == ( + 7, + workspace, + ) + assert calls == [layouts] * 4 + + @pytest.mark.parametrize("reference", ["absent", "expired", "smaller"]) def test_hybridep_high_water_needs_a_live_larger_graph( hybrid_checkpoint_rank, reference From e204ee84912df2a39f7c93de5c8f402212d5dc37 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:49:23 +0000 Subject: [PATCH 18/19] Refuse CP layout facts live pricing cannot produce Validate an empty GDN layout's length, require the TP1/CP2/PP1 topology that layout pricing is limited to, and refuse a recorded group_layouts argument in any estimate. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_misses.py | 1 + src/art/trainer_rank/_planner_replay.py | 12 ++++++++--- tests/unit/test_grouped_planner_replay.py | 25 ++++++++++++++++++++++- 3 files changed, 34 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index 028c87f1e..46fa80d4e 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -723,6 +723,7 @@ def replay( ) ) or arguments.get("slot_refs") + or arguments.get("group_layouts") or arguments.get("head_workspace_bytes", 0) or any(arguments.get("checkpoint_floor", (0, 0))) ): diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index e5a7b7a04..7ef1b4a92 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -492,8 +492,10 @@ def integer(value: Any, *, minimum: int = 0) -> None: not group["grad"] or type(ranks) is not list or not 0 < len(ranks) <= 64 - or (gdn_rows is not None and type(gdn_rows) is not list) - or len(gdn_rows or ranks) != len(ranks) + or ( + gdn_rows is not None + and (type(gdn_rows) is not list or len(gdn_rows) != len(ranks)) + ) or type(layout["attention_retained"]) is not list or len(layout["attention_retained"]) != len(ranks) ): @@ -823,7 +825,11 @@ def runtime_arguments( raise ValueError("head backward staging disagrees with recorded facts") layouts = None if groups[0]["layout"] is not None: - if len(groups[0]["layout"]["attention_rows"]) != self._topology_key()[2]: + # Live layout pricing is CP2 at TP1/PP1 only (_layout_pricing_supported). + _, tp, cp, pp = self._topology_key() + if (tp, cp, pp) != (1, 2, 1) or len( + groups[0]["layout"]["attention_rows"] + ) != 2: raise ValueError("CP layout facts disagree with the recorded topology") layouts = tuple( _impl._GroupLayout( diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 80e799ef5..10d31441f 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -777,6 +777,7 @@ def changed(edit): "inputs_type", "ranks", "argument", + "empty_gdn", ], ) def test_cp_layout_fact_validation_rejects_forged_input(change, monkeypatch, tmp_path): @@ -805,13 +806,35 @@ def test_cp_layout_fact_validation_rejects_forged_input(change, monkeypatch, tmp for key in ("attention_rows", "gdn_rows", "attention_retained"): layout[key].append(1) message = "recorded topology" - else: + elif change == "argument": item["arguments"]["group_layouts"] = [dict(layout)] message = "immutable runtime" + else: + layout["gdn_rows"] = [] + message = "invalid CP layout facts" with pytest.raises(ValueError, match=message): reports.replay(report) +def test_cp_layouts_need_the_live_cp2_topology(tmp_path): + # Layout pricing is CP2-only: layouts forged onto a CP1 report are refused. + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + (group,) = facts["groups"] + group["layout"] = { + "attention_rows": [group["rows"]], + "gdn_rows": None, + "attention_retained": [0], + } + facts["layer_gdn_inputs"] = [False] * facts["checkpoint_layers"] + with pytest.raises(ValueError, match="recorded topology"): + reports.replay(report) + + @pytest.mark.parametrize("name", ["_plan_group_layouts", "_layer_gdn_inputs"]) def test_custom_layout_reader_is_explicitly_incomplete(name, monkeypatch): from art.trainer_rank import _planner_replay From bf79e4468d6ec801185d6857a3945cfa7db9de7e Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 07:10:53 +0000 Subject: [PATCH 19/19] Test that an ungrouped estimate cannot record CP layouts Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_planner_reports.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index a9b26b6af..f834c598e 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -384,6 +384,13 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): ] = rows with pytest.raises(ValueError, match="immutable runtime"): reports.replay(changed) + # CP layouts are grouped runtime facts; no estimate records them itself. + changed = json.loads(path.read_bytes()) + changed["replay"]["memory_replay"]["estimates"][0]["arguments"]["group_layouts"] = [ + {"attention_rows": [1, 1]} + ] + with pytest.raises(ValueError, match="immutable runtime"): + reports.replay(changed) for field in ( "checkpoint_input_gradient", "checkpoint_workspace",