diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index cce655bd8..c3bc7ba1f 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -246,6 +246,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 \ @@ -297,4 +299,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 f581ae9b4..46fa80d4e 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -716,12 +716,14 @@ def replay( name in arguments for name in ( "slot_refs", + "group_layouts", "head_workspace_bytes", "checkpoint_floor", ) ) ) 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 2c9d9e747..7ef1b4a92 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -137,6 +137,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", # Producers of recorded arguments replay checks against the facts. "_plan_group_routed_rows", "_plan_head_backward_traced", @@ -191,8 +201,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) @@ -271,6 +283,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: @@ -291,6 +314,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 { @@ -313,8 +337,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, @@ -323,6 +352,7 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: name: getattr(rank, "_" + name)() if layers else None for name in _RECOMPUTE_READERS }, + "layer_gdn_inputs": list(inputs), "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -354,12 +384,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", @@ -398,6 +429,7 @@ def integer(value: Any, *, minimum: int = 0) -> None: "head_target_rows", "adapter", "moe_covered", + "layout", "gdn", }, ) @@ -451,6 +483,26 @@ 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) != 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"}) @@ -505,6 +557,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") @@ -680,10 +749,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, @@ -692,8 +762,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) @@ -747,6 +823,24 @@ def runtime_arguments( and facts["head_vocabulary"] ): raise ValueError("head backward staging disagrees with recorded facts") + layouts = None + if groups[0]["layout"] is not None: + # 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( + 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"]), @@ -779,6 +873,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 d3540efe4..10d31441f 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -720,3 +720,128 @@ 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", + "empty_gdn", + ], +) +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" + 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 + + 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 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",