diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 924abc0cc..cce655bd8 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -233,9 +233,11 @@ jobs: tests/unit/test_trainer_rank_profile_warm.py \ tests/unit/test_trainer_rank_tp_floor.py \ tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \ + tests/unit/test_trainer_rank_adapter_gradient_memory.py \ tests/unit/test_trainer_rank_slot_memory.py \ tests/unit/test_trainer_rank_moe_memory.py \ tests/unit/test_trainer_rank_head_memory.py \ + tests/unit/test_trainer_rank_head_stage_memory.py \ tests/unit/test_grouped_planner_replay.py \ tests/unit/test_planner_runtime_fact_guards.py \ tests/unit/test_planner_replay_owner_budget.py \ @@ -282,9 +284,11 @@ jobs: --ignore=tests/unit/test_trainer_rank_profile_warm.py \ --ignore=tests/unit/test_trainer_rank_tp_floor.py \ --ignore=tests/unit/test_trainer_rank_checkpoint_gradient_memory.py \ + --ignore=tests/unit/test_trainer_rank_adapter_gradient_memory.py \ --ignore=tests/unit/test_trainer_rank_slot_memory.py \ --ignore=tests/unit/test_trainer_rank_moe_memory.py \ --ignore=tests/unit/test_trainer_rank_head_memory.py \ + --ignore=tests/unit/test_trainer_rank_head_stage_memory.py \ --ignore=tests/unit/test_grouped_planner_replay.py \ --ignore=tests/unit/test_planner_runtime_fact_guards.py \ --ignore=tests/unit/test_planner_replay_owner_budget.py \ diff --git a/src/art/megatron/context_parallel/runtime.py b/src/art/megatron/context_parallel/runtime.py index 821b60e76..531cb8829 100644 --- a/src/art/megatron/context_parallel/runtime.py +++ b/src/art/megatron/context_parallel/runtime.py @@ -433,6 +433,45 @@ def context_parallel_rank_model_token_counts( ) +def context_parallel_model_token_total( + *, + 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, +) -> int: + """Return the CP group's model rows in its larger physical layout. + + Dispatch runs at least one row on every rank; an empty rank's padding row + passes through the model too. + """ + 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, + ) + ) + total = sum( + max(1, count) for count in bundle.token_layout_index.token_counts_by_rank + ) + if not build_gdn_execution_spec: + return total + decision = _plan_gdn_global_execution( + planning_key=planning_key, + bundle=bundle, + topology=topology, + gdn_planner_config=gdn_planner_config, + ) + return max(total, sum(max(1, count) for count in decision.gdn_token_counts_by_rank)) + + def _normalized_chunk_size( *, valid_tokens: int, diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index a25505c53..2ec11b766 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -17,6 +17,7 @@ from dataclasses import dataclass, replace from dataclasses import field as dataclass_field from functools import partial +import json import logging import math import os @@ -124,6 +125,34 @@ class TopK: _MEMORY_SAFETY_FACTOR = 1.10 _MEMORY_RESERVE_FRACTION = 0.03 _HEAD_CHUNK_TOKENS = 512 +# An unprofiled full-recompute gradient wave's first execution keeps two fixed +# 32 MiB transients live at its peak (Qwen3.6-35B-A3B CP2: the RoPE frequencies +# and a frozen linear's output, at 2k to 20k tokens); warm waves do not. +_COLD_RECOMPUTE_TRANSIENT_BYTES = 64 * 2**20 + +# Per local row, index state the backward keeps beside the RoPE embedding and +# hidden-width tensors: int64 positions and row maps, CP block masks and GDN +# exchange plans. Qwen3.6-35B-A3B CP2 traces at the head's backward peak: +# 106-181 bytes per row at 3.5k-8.7k rows. +_BACKWARD_ROW_STATE_BYTES = 256 + +# A head chunk below the fused statistics' row minimum takes the FP32 fallback +# (_vocab_parallel_log_z): the BF16 logits, their FP32 copy, the shifted copy +# and its saved exponent in the recompute, then the exponent's gradient, its +# BF16 cast and the target gather's gradient in backward. At most this many +# BF16 logits-sized buffers of that chunk (about 18 bytes per logit); derived +# from the code, not traced. +_HEAD_FALLBACK_BUFFERS = 9 + +# Which fused head statistics kernels have run in this process, and whether any +# call fell back to FP32 after an error: staging trusts only a proven path. +_TRITON_STATS_STATE: dict[str, Any] = {"succeeded": set(), "failed": False} + +# Set while a gradient group of a plan priced with a staged head projects it; +# per thread, and captured into each chunk's checkpoint for its recompute. +_HEAD_STATISTICS_STRICT: ContextVar[bool] = ContextVar( + "trainer_rank_head_statistics_strict", default=False +) _PLANNER_REFINEMENT_BUDGET = 2_000 _LAYOUT_SELECTION_CACHE_LIMIT = 64 @@ -1060,6 +1089,13 @@ class _SubforwardCost: # HybridEP buffer growth before the safety factor. It is in required, not # retained, and persists across a split, which charges the largest once. hybridep_growth: int = 0 + # Adapter gradients the recompute backward holds beyond the boundaries it + # has released (``_checkpoint_adapter_gradient_bytes``), before the safety + # factor. Split children training the same slots share them, so a split + # charges the largest once; ``..._slots`` names those slots (sorted JSON + # of kind/name pairs, "" when none), keeping the cost JSON-serializable. + checkpoint_adapter_gradient: int = 0 + checkpoint_adapter_gradient_slots: str = "" @property def ephemeral(self) -> int: @@ -1541,8 +1577,30 @@ def _expert_lora_weight_storage( return (transposes if a.shape[2] < 8 else 0, transposes, effective) -# Routed rows per rank at EP>1, relative to balanced routing (see below). -_EP_ROUTED_ROW_ALLOWANCE = 1.5 +# Routed rows on the most loaded rank at EP>1, relative to its balanced share. +# Expert-shard load is uneven per layer, from the router's expert preferences, +# and larger batches do not average it away. Qwen3.6-35B-A3B on 3.5M tokens of +# retail agent trajectories, worst layer in 200k-token batches at EP2 / EP4 / +# EP8: pretrained up to 1.22 / 1.40 / 1.62, a trained policy up to 1.24 / 1.41 / +# 1.61; a small rollout sample reached 1.95 at EP8. One production EP2 run was +# inferred at 1.35. These samples bound what was measured, not all routing. +# Unmeasured EP sizes use the next measured one; above EP8 the allowance grows +# with log2(EP) up to EP itself (every pair on one rank). +_EP_ROUTED_ROW_ALLOWANCE = {2: 1.4, 4: 1.6, 8: 2.0} + + +def _ep_routed_row_allowance(ep: int) -> float: + if ep <= 1: + return 1.0 + for size, allowance in sorted(_EP_ROUTED_ROW_ALLOWANCE.items()): + if ep <= size: + return allowance + return min(float(ep), 2.0 + 0.4 * math.log2(ep / 8)) + + +# Transformer Engine's Hopper cuBLAS workspaces: one per grouped-GEMM stream +# (four) plus the plain GEMM's, each 32 MiB + 1 KiB. +_TE_CUBLAS_WORKSPACE_BYTES = 5 * (32 * 2**20 + 1024) def _moe_dispatcher_supported( @@ -1559,6 +1617,29 @@ def _moe_dispatcher_supported( ) +def _ep_group_is_cp_group(shape: ParallelShape) -> bool: + """Whether this rank's expert-parallel group is exactly its CP group. + + HybridEP then dispatches that CP group's rows across it: at balanced + routing each rank receives the group's rows over EP, however CP split them. + """ + if shape.ep <= 1 or shape.ep != shape.cp or (shape.tp, shape.etp) != (1, 1): + return False + if not dist.is_available() or not dist.is_initialized(): + return False + try: + from megatron.core import parallel_state as ps + except ModuleNotFoundError: + return False + expert = ps.get_expert_model_parallel_group(check_initialized=False) + context = ps.get_context_parallel_group(check_initialized=False) + if expert is None or context is None: + return False + return sorted(dist.get_process_group_ranks(expert)) == sorted( + dist.get_process_group_ranks(context) + ) + + def _hybridep_rows_per_rank(capacity: int, ranks: int) -> int: """HybridEP's allocated rows per rank: TMA-aligned, at least 512, padded to the 64-row combine chunk.""" @@ -1590,8 +1671,15 @@ def _moe_output_bytes_per_token( checkpoint_grad: bool = False, converted_stages: list[tuple[int, int]] | None = None, slot_ref: "LoRASlotRef | None" = None, + shared_bytes: list[int] | None = None, + enclosed: list[bool] | None = None, ) -> int: - """Known routed-expert working set, not a complete model/compiled bound.""" + """Known routed-expert working set, not a complete model/compiled bound. + + ``shared_bytes`` collects each layer's shared-expert part of the per-token + coefficient and stages; that part follows local rows, not routed rows. + ``enclosed`` records, per MoE layer, whether its FC1 stage is priced too. + """ # CP shards rows, not the per-token working set. At EP>1 only ART's # HybridEP flex dispatcher is modeled; TP and ETP are not. if (shape.tp, shape.etp) != (1, 1): @@ -1612,9 +1700,12 @@ def _moe_output_bytes_per_token( from art.megatron.lora import LoRA, MLPExpertsLinearFC1LoRA, MLPExpertsLinearFC2LoRA # HybridEP hands each rank the pairs routed to its local experts, already - # permuted. Balanced routing gives local tokens x top-k, as at EP1; a - # pretrained CP2/EP2 run put about 1.35x that on one rank. - routed_allowance = _EP_ROUTED_ROW_ALLOWANCE if shape.ep > 1 else 1 + # permuted. Balanced routing gives local tokens x top-k, as at EP1. + routed_allowance = _ep_routed_row_allowance(shape.ep) + # Routed H-wide inputs held at the expert stage. The EP1 all-to-all path + # keeps its permuted rows and their expert-sorted copy; HybridEP permutes + # while it dispatches and returns one tensor. + dispatched = 1 if shape.ep > 1 else 2 coefficient = 0 for chunk in model: for layer in chunk.modules(): @@ -1708,8 +1799,6 @@ def _moe_output_bytes_per_token( and fc1.fused_gate_up and not fc1.non_gated and fc1.out_features == 2 * inputs.shape[-2] - # HybridEP keeps one dispatched H-wide input where the - # EP1 all-to-all keeps two; charging two is conservative. and getattr(dispatcher, "ep_size", None) == shape.ep and getattr(dispatcher, "tp_size", None) == 1 and getattr(dispatcher, "num_local_experts", 0) > 1 @@ -1718,10 +1807,10 @@ def _moe_output_bytes_per_token( and getattr(experts, "offload_moe_act", None) is False and getattr(experts, "activation_recompute", None) is False ): - # The two dispatched H-wide inputs and FC1 gate/up sum - # remain live at the FC2 sum, including in the observed - # compiled path. This is one stage, not a backward bound. - features += 2 * fc2.out_features + fc1.out_features + # The dispatched H-wide inputs and FC1 gate/up sum remain + # live at the FC2 sum, including in the observed compiled + # path. This is one stage, not a backward bound. + features += dispatched * fc2.out_features + fc1.out_features enclosing_fc1 = fc1 shared = _shared_expert_output_bytes_per_token(layer) if ( @@ -1733,10 +1822,15 @@ def _moe_output_bytes_per_token( # Gate-score backward saves a distinct pre-gate X. Charge it # beside this layer's returned X, not another layer's maximum. shared += shared - routed_rows = math.ceil(config.moe_router_topk * routed_allowance) - row_bytes = routed_rows * features * weights.element_size() + shared + if shared_bytes is not None: + shared_bytes.append(shared) + routed_rows = config.moe_router_topk * routed_allowance + row_bytes = ( + math.ceil(routed_rows * features * weights.element_size()) + shared + ) coefficient = max(coefficient, row_bytes) storage = _expert_lora_weight_storage(lora, slot_ref) + fc1_stages = False if converted_stages is not None and storage is not None: padded, transposes, effective = storage saved_fc1, rank_fc1 = 0, 0 @@ -1765,13 +1859,14 @@ def _moe_output_bytes_per_token( and first_tensors[1].shape[2] == enclosing_fc1.out_features ): first_padding, first_transposes, first_rank = first - # FC1 retains both routed H inputs and its base O1 + fc1_stages = True + # FC1 retains the routed H inputs and its base O1 # while producing adapter O1. Its sum is not live yet. converted_stages.append( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 2 * enclosing_fc1.out_features + first_rank ) @@ -1785,7 +1880,7 @@ def _moe_output_bytes_per_token( ( routed_size * ( - 2 * fc2.out_features + dispatched * fc2.out_features + 3 * enclosing_fc1.out_features + (first_rank if checkpoint_grad else 0) ) @@ -1839,6 +1934,23 @@ def _moe_output_bytes_per_token( + 2 * (experts_count + 1) * 4, ) ) + if enclosed is not None: + # FC1 is covered when priced beside FC2 and its converted + # stages are priced too, unless it has no adapter or the + # selected slot has no FC1 tensors to convert. + adapter = getattr(enclosing_fc1, "lora", None) + inactive = adapter is None or ( + slot_ref is not None + and type(adapter) is LoRA + and "_slot" not in vars(adapter) + and _slot_lora_tensors(adapter, slot_ref) is None + ) + enclosed.append(enclosing_fc1 is not None and (fc1_stages or inactive)) + if converted_stages is not None: + # The EP allowance gives fractional routed rows; round each stage up. + converted_stages[:] = [ + (math.ceil(per_row), fixed) for per_row, fixed in converted_stages + ] return coefficient @@ -1964,9 +2076,15 @@ def memory_field(name: str, default: Any = None) -> Any: ) forward_stages: list[tuple[int, int]] = [] gradient_stages: list[tuple[int, int]] = [] + forward_shared: list[int] = [] + gradient_shared: list[int] = [] + gradient_enclosed: list[bool] = [] self._moe_output_bytes_per_token = ( _moe_output_bytes_per_token( - runtime.model, self._parallel_shape, converted_stages=forward_stages + runtime.model, + self._parallel_shape, + converted_stages=forward_stages, + shared_bytes=forward_shared, ) if self._moe_layers else 0 @@ -1980,6 +2098,8 @@ def memory_field(name: str, default: Any = None) -> Any: self._parallel_shape, checkpoint_grad=True, converted_stages=gradient_stages, + shared_bytes=gradient_shared, + enclosed=gradient_enclosed, ) if self._moe_layers else 0 @@ -1991,6 +2111,26 @@ def memory_field(name: str, default: Any = None) -> Any: self._moe_gradient_stages = ( tuple(gradient_stages) if self._moe_checkpoint_grad_bytes_per_token else () ) + self._moe_forward_shared_bytes = ( + max(forward_shared, default=0) if self._moe_output_bytes_per_token else 0 + ) + self._moe_gradient_shared_bytes = ( + max(gradient_shared, default=0) + if self._moe_checkpoint_grad_bytes_per_token + else 0 + ) + # Recompute is covered only if every decoder layer is a priced MoE + # layer whose FC1 stage is enclosed; otherwise dense MLP or FC1 work + # the floor does not price keeps the per-boundary gradient allowance. + self._moe_gradient_enclosed = ( + tuple(gradient_enclosed) + if self._moe_checkpoint_grad_bytes_per_token + else () + ) + self._moe_recompute_covered = len( + self._moe_gradient_enclosed + ) == self._num_layers and all(self._moe_gradient_enclosed) + self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, device_memory_bytes=device_memory, @@ -2900,14 +3040,16 @@ def _subforward_cost( logical_tokens: int, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, hybridep_growth_bytes: int = 0, ) -> _SubforwardCost: checkpoint_memory = self._checkpoint_memory_floor( - group_rows, slot_refs, gdn_segments + group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows ) required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, @@ -2916,6 +3058,7 @@ def _subforward_cost( logical_tokens=logical_tokens, gdn_segments=gdn_segments, group_rows=group_rows, + group_routed_rows=group_routed_rows, slot_refs=slot_refs, head_workspace_bytes=head_workspace_bytes, checkpoint_floor=checkpoint_floor, @@ -2935,25 +3078,45 @@ def _subforward_cost( checkpoint_floor[0], ), ) - # One logical BF16 input gradient per eligible full/uniform/1 boundary. - # This partial peak allowance is not evidence of simultaneous distinct - # backing stores, nor a bound for compiler saves or other backward work. - # Keep it out of forward retention, including the cold fallback above. - gradient = checkpoint_retained + # Input gradients live at the recomputed layer's peak; kept out of + # forward retention, including the cold fallback above. + gradient = self._checkpoint_input_gradient_bytes( + group_rows, slot_refs, retained=checkpoint_retained + ) + gradient_slots = self._gradient_slots(group_rows, slot_refs) + adapter_gradient = ( + self._checkpoint_adapter_gradient_bytes( + self._checkpoint_gradient_groups(group_rows, slot_refs) + ) + if gradient + else 0 + ) checkpoint_retained = output_bytes + max( checkpoint_retained, checkpoint_floor[0] ) - checkpoint_workspace = max( - checkpoint_workspace, head_workspace_bytes, checkpoint_floor[1] - ) + decoder_workspace = max(checkpoint_workspace, checkpoint_floor[1]) + checkpoint_workspace = max(decoder_workspace, head_workspace_bytes) forward_required = required if gradient: + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + peak = checkpoint_workspace + adapter_gradient + if head_stage is not None: + # A split adds its children's adapter gradients to the largest + # workspace, and one child's head can follow another's decoder + # backward: keep the whole head stage there. + checkpoint_workspace = max(decoder_workspace, head_stage) + # An untraced head's buffers can exceed head_workspace_bytes: + # it keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) + if self._memory_profiles.get(signature) is None: + checkpoint_workspace += _COLD_RECOMPUTE_TRANSIENT_BYTES + peak += _COLD_RECOMPUTE_TRANSIENT_BYTES required = max( required, - int( - (checkpoint_retained + checkpoint_workspace + gradient) - * _MEMORY_SAFETY_FACTOR - ), + int((checkpoint_retained + gradient + peak) * _MEMORY_SAFETY_FACTOR), ) # HybridEP buffer growth stays allocated through the forward and # backward peaks, but is not forward retention. @@ -2965,6 +3128,22 @@ def _subforward_cost( checkpoint_input_gradient=gradient, checkpoint_peak_increment=required - forward_required, hybridep_growth=hybridep_growth_bytes, + checkpoint_adapter_gradient=adapter_gradient, + checkpoint_adapter_gradient_slots=json.dumps( + [ + [ref.kind, ref.name] + for ref in sorted( + gradient_slots, + key=lambda ref: ( + ref.kind, + ref.name is not None, + ref.name or "", + ), + ) + ] + ) + if adapter_gradient + else "", ) def last_forward_telemetry(self) -> dict[str, Any]: @@ -3378,6 +3557,8 @@ def _telemetry_signature(cls, plan: _AnyForwardPlan) -> dict[str, object]: } def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: + # A head priced as staged must not silently widen (_head_backward_traced). + staged = bool(getattr(plan, "_head_staged", False)) outputs = [ ForwardOutput(None, None, None, None, checkpoint, no_grad) for checkpoint, no_grad in plan.output_metadata @@ -3399,7 +3580,13 @@ def _execute_flat_plan(self, plan: _FlatForwardPlan) -> list[AnyForwardOutput]: with torch.set_grad_enabled(group.grad_enabled): with use_lora_slot(group.slot_ref): prepared = self._prepare_packed_forward(group.packed) - item_outputs = self._forward_packed(group.items, prepared) + strict = _HEAD_STATISTICS_STRICT.set( + staged and group.grad_enabled + ) + try: + item_outputs = self._forward_packed(group.items, prepared) + finally: + _HEAD_STATISTICS_STRICT.reset(strict) item_outputs = [ replace( output, @@ -4243,6 +4430,7 @@ def _project_vocab_parallel( from torch.utils.checkpoint import checkpoint model = _language_model(self.runtime.model[0]) + strict_statistics = _HEAD_STATISTICS_STRICT.get() max_top_k = max((int(item.request.top_k or 0) for item in items), default=0) need_log_z = any( item.labels is not None or item.request.top_k is not None for item in items @@ -4260,6 +4448,8 @@ def _project_vocab_parallel( output_weight=output_weight, need_log_z=need_log_z, max_top_k=max_top_k, + # Captured now: the backward recompute keeps this plan's mode. + strict_statistics=strict_statistics, use_reentrant=False, ) logit_start, logit_end = logit_bounds[chunk_index : chunk_index + 2] @@ -4340,6 +4530,7 @@ def _local_head_stats( output_weight: torch.Tensor | None, need_log_z: bool, max_top_k: int, + strict_statistics: bool = False, ) -> tuple[ torch.Tensor, torch.Tensor | None, @@ -4353,11 +4544,17 @@ def _local_head_stats( log_z: torch.Tensor | None = None local_topk: tuple[torch.Tensor, torch.Tensor] | None = None if need_log_z: - topk_stats = _try_triton_local_topk_stats(local_logits, k=max_top_k) + topk_stats = _try_triton_local_topk_stats( + local_logits, k=max_top_k, strict=strict_statistics + ) logsumexp_stats = ( cast( tuple[torch.Tensor, torch.Tensor] | None, - _try_triton_stats("local_logsumexp_stats", local_logits), + _try_triton_stats( + "local_logsumexp_stats", + local_logits, + strict=strict_statistics, + ), ) if topk_stats is None else None @@ -4698,12 +4895,66 @@ def _gather_tensor_parallel_logits(self, logits: torch.Tensor) -> torch.Tensor: _record_split_memory_floor = _memory._record_split_memory_floor _split_plan_memory_check = _memory._split_plan_memory_check _head_workspace_bytes = _memory._head_workspace_bytes + + def _cp_group_model_tokens( + self, + batch: PrefixTreePack, + *, + topology: "ParallelTopology", + ) -> int: + """The CP group's model rows in its larger physical layout.""" + from art.megatron.context_parallel.runtime import ( + context_parallel_model_token_total, + ) + from art.megatron.training.microbatches import ( + _context_parallel_config_for_provider, + _gdn_planner_config_for_provider, + ) + + handler = self.runtime.model_support_handler + return context_parallel_model_token_total( + group_ids=batch.group_ids, + parent_ids=batch.parent_ids, + topology=topology, + config=_context_parallel_config_for_provider( + self.runtime.provider, + self.device, + handler, + ), + 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 + ), + ) + _group_head_workspace_bytes = _memory._group_head_workspace_bytes + _triton_min_rows = _memory._triton_min_rows + _head_backward_traced = _memory._head_backward_traced + _te_workspace_growth_bytes = _memory._te_workspace_growth_bytes + _moe_checkpoint_state_bytes_per_token = ( + _memory._moe_checkpoint_state_bytes_per_token + ) + _moe_recompute_covered_for = _memory._moe_recompute_covered_for + _checkpoint_input_gradient_bytes = _memory._checkpoint_input_gradient_bytes + _checkpoint_gradient_covered = _memory._checkpoint_gradient_covered + _checkpoint_head_stage_bytes = _memory._checkpoint_head_stage_bytes + _backward_row_state_bytes = _memory._backward_row_state_bytes + _mixer_activation_widths = _memory._mixer_activation_widths + _recomputed_mixer_bytes_per_token = _memory._recomputed_mixer_bytes_per_token + _adapter_gradient_head = staticmethod(_memory._adapter_gradient_head) + _plan_head_backward_traced = _micro_batch_planner._plan_head_backward_traced + _plan_group_routed_rows = _micro_batch_planner._plan_group_routed_rows _plan_head_workspace_bytes = _memory._plan_head_workspace_bytes _plan_hybridep_growth_bytes = _memory._plan_hybridep_growth_bytes _checkpoint_moe_bytes_per_token = _memory._checkpoint_moe_bytes_per_token _moe_workspace_bytes = _memory._moe_workspace_bytes _checkpoint_memory_floor = _memory._checkpoint_memory_floor + _gradient_slots = staticmethod(_memory._gradient_slots) + _pending_adapter_gradient_bytes = _memory._pending_adapter_gradient_bytes + _checkpoint_gradient_groups = _memory._checkpoint_gradient_groups + _checkpoint_adapter_gradient_bytes = _memory._checkpoint_adapter_gradient_bytes + _adapter_gradient_walk = staticmethod(_memory._adapter_gradient_walk) _retained_memory_bytes = _memory._retained_memory_bytes _estimate_flat_forward = _memory._estimate_flat_forward _update_peak_memory_profile = _memory._update_peak_memory_profile @@ -4846,6 +5097,18 @@ def _active_logical_tokens(requests: Sequence[AnyForwardInput]) -> int: _PACKED_PRICED_MIN_REQUEST_TOKENS = 64 +def _traced_states(traced: Sequence[bool | None]) -> tuple[bool, ...]: + """Head stagings a plan's gradient groups allow (``_head_backward_traced``). + + Staged only when every group is traced; both when any is undecided. + """ + if not traced or any(state is False for state in traced): + return (False,) + if all(state is True for state in traced): + return (True,) + return (True, False) + + def _packed_priced(signature: "_MemorySignature", one_layer_recompute: bool) -> bool: # Measured only under one-layer full recompute: other recompute modes keep # more GDN states and activations live per segment and logical row. @@ -5380,6 +5643,7 @@ def _try_triton_local_topk_stats( local_logits: torch.Tensor, *, k: int, + strict: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None: if k <= 0 or k > int( os.environ.get("ART_TRAINER_RANK_TRITON_FUSED_TOPK_MAX", "10") @@ -5390,6 +5654,7 @@ def _try_triton_local_topk_stats( _try_triton_stats( "local_topk_stats", local_logits, + strict=strict, k=min(k, int(local_logits.shape[1])), ), ) @@ -5398,8 +5663,16 @@ def _try_triton_local_topk_stats( def _try_triton_stats( name: str, local_logits: torch.Tensor, + *, + strict: bool = False, **kwargs: object, ) -> object | None: + """The fused statistics, or None for the FP32 fallback. + + ``strict``: a plan whose price relied on them raises instead of falling + back after an error (``_head_backward_traced``). Too few rows still fall + back: the head stage prices that (``_HEAD_FALLBACK_BUFFERS``). + """ if not local_logits.is_cuda: return None if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() in { @@ -5412,11 +5685,19 @@ def _try_triton_stats( try: from art.trainer_rank import topk - return getattr(topk, name)(local_logits, **kwargs) - except Exception: + result = getattr(topk, name)(local_logits, **kwargs) + except Exception as error: + _TRITON_STATS_STATE["failed"] = True + if strict: + raise RuntimeError( + "Fused head statistics failed in a plan admitted on their " + "memory; the FP32 fallback would exceed its price" + ) from error if os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() == "strict": raise return None + _TRITON_STATS_STATE["succeeded"].add(name) + return result def _vocab_parallel_topk_from_local( diff --git a/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index 70e3c9250..6a4654222 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, Sequence +from collections.abc import Iterable, Iterator, Sequence from contextlib import nullcontext import hashlib import math @@ -45,12 +45,32 @@ def _split_required_memory(costs: Sequence[_impl._SubforwardCost]) -> int: if any(cost.checkpoint_input_gradient for cost in costs): # The caller owns all returned graphs. A calibrated forward-retained # discount cannot replace the sum of their input-gradient extents. + # Children training the same slots share their adapter gradients, + # allocated once by whichever child's backward reaches a layer + # first; the largest child's extra covers any order. Different + # slots have disjoint gradients, charged per child. + adapter = [ + cost.checkpoint_adapter_gradient + for cost in costs + if cost.checkpoint_adapter_gradient + ] + shared = ( + len( + { + cost.checkpoint_adapter_gradient_slots + for cost in costs + if cost.checkpoint_adapter_gradient + } + ) + <= 1 + ) checkpoint = ( sum( cost.checkpoint_retained + cost.checkpoint_input_gradient for cost in costs ) + max(cost.checkpoint_workspace for cost in costs) + + ((max(adapter) if shared else sum(adapter)) if adapter else 0) + growth ) required = max(required, int(checkpoint * _impl._MEMORY_SAFETY_FACTOR)) @@ -418,13 +438,14 @@ def _moe_workspace_terms( *, checkpoint_grad: bool = False, slot_ref: "LoRASlotRef | None" = None, -) -> tuple[int, tuple[tuple[int, int], ...]]: +) -> tuple[int, tuple[tuple[int, int], ...], int]: """Maximum of same-layer affine stages, not a retained multi-layer bank. The constructor cache covers original tensors. Explicit slots are repriced from their tensor metadata and original owners, including this rank's exact dispatcher wrapper. Ordinary non-checkpoint gradients - retain only forward-stage coverage. + retain only forward-stage coverage. The third term is the shared + expert's part of the coefficient, which stays on the local rows. """ coefficient = ( self._checkpoint_moe_bytes_per_token() @@ -436,8 +457,16 @@ def _moe_workspace_terms( "_moe_gradient_stages" if checkpoint_grad else "_moe_forward_stages", (), ) + shared = getattr( + self, + "_moe_gradient_shared_bytes" + if checkpoint_grad + else "_moe_forward_shared_bytes", + 0, + ) if slot_ref is not None and slot_ref.name is not None: selected: list[tuple[int, int]] = [] + slot_shared: list[int] = [] coefficient = ( _impl._moe_output_bytes_per_token( self.runtime.model, @@ -445,11 +474,13 @@ def _moe_workspace_terms( checkpoint_grad=checkpoint_grad, converted_stages=selected, slot_ref=slot_ref, + shared_bytes=slot_shared, ) if self._moe_layers else 0 ) stages = tuple(selected) if coefficient else () + shared = max(slot_shared, default=0) if coefficient else 0 if type(stages) is not tuple or any( type(stage) is not tuple or len(stage) != 2 @@ -457,20 +488,27 @@ def _moe_workspace_terms( for stage in stages ): raise ValueError("Invalid constructor converted-weight stages") - return coefficient, stages + if type(shared) is not int or not 0 <= shared <= coefficient: + raise ValueError("Invalid constructor shared-expert coefficient") + return coefficient, stages, shared def _moe_workspace_from_terms( - rows: int, terms: tuple[int, tuple[tuple[int, int], ...]] + rows: int, + terms: tuple[int, tuple[tuple[int, int], ...], int], + routed_rows: int | None = None, ) -> int: - coefficient, stages = terms - return ( + """``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 + ( max( - rows * coefficient, - *(rows * per_row + fixed for per_row, fixed in stages), + routed * coefficient, + *(routed * per_row + fixed for per_row, fixed in stages), ) if stages and rows > 0 - else rows * coefficient + else routed * coefficient ) @@ -478,12 +516,14 @@ def _moe_workspace_bytes( self: TrainerRank, rows: int, *, + routed_rows: int | None = None, checkpoint_grad: bool = False, slot_ref: "LoRASlotRef | None" = None, ) -> int: return _moe_workspace_from_terms( rows, _moe_workspace_terms(self, checkpoint_grad=checkpoint_grad, slot_ref=slot_ref), + routed_rows, ) @@ -567,26 +607,41 @@ def _checkpoint_memory_floor( group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None = None, gdn_segments: int = 0, + routed_rows: tuple[int, ...] | 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``. + """ 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 + self, group_rows, slot_refs, gdn_segments, layers, routed_rows ) - # HybridEP runtime state is intentionally outside grouped CPU replay v1. - hybrid_rows = None + # HybridEP runtime state is intentionally outside grouped CPU replay. if ( any(grad for _, grad in group_rows) and self._topology_key() == (1, 1, 2, 1) and self._parallel_shape == _impl.ParallelShape(tp=1, cp=2, ep=2, etp=1) and self._moe_memory_supported ): - hybrid_rows = max(rows for rows, _ in group_rows) + # Recompute runs after _execute_flat_plan restores the communication + # high-water. Combine allocates a fresh BF16 [P, H] before cropping; + # this is separate from already-held native buffer capacity. Do not + # prune graph references or reset execution state while estimating. + rows = max(rows for rows, _ in group_rows) if any(ref() is not None for ref in self._pending_hybridep_graphs): - hybrid_rows = max(hybrid_rows, self._hybridep_rows_high_water) - if hybrid_rows is not None: - workspace = max(workspace, -(-hybrid_rows // 4) * 4 * self._hidden_size * 2) + rows = max(rows, self._hybridep_rows_high_water) + # The combine output and the TE workspaces are live together. + workspace = max( + workspace, + -(-rows // 4) * 4 * self._hidden_size * 2 + + self._te_workspace_growth_bytes(), + ) return retained, workspace @@ -596,35 +651,591 @@ def _checkpoint_floor_from_facts( slot_refs: tuple["LoRASlotRef | None", ...] | None, gdn_segments: int, layers: int, + routed_rows: tuple[int, ...] | 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 + enter the term; already-live graphs remain in the availability baseline. + The workspace is one MoE stage plus what else is live beside it. + Gradient groups recompute the layer, so its attention or GDN mixer + keeps its saved activations across the MoE stage. No-grad groups keep + decoder input, current layer input, its MLP residual and norm output. + Count these four row tensors separately from returned outputs, allowing + storage aliases. This is not a bound for custom preprocessing or all of + backward. With sequence parallelism a rank saves only its shard of each + boundary; that is priced only where ``_sequence_parallel_floor_covered`` + holds, and there, for gradient waves, the recomputed GDN layer's + recurrent states for ``gdn_segments`` (gradient groups' segments) plus + padding, as traced (the recomputed mixer is not added at TP > 1). + """ if not layers: return 0, 0 gradient_rows = sum(rows for rows, grad in group_rows if grad) _, tp, _, _ = self._topology_key() - # Physical rows are padded to a multiple of TP; each rank saves its shard. - retained = ( - sum(-(-rows // tp) for rows, grad in group_rows if grad) - * layers - * self._hidden_size - * 2 - ) - if gradient_rows: - self._checkpoint_moe_bytes_per_token() refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + if tp > 1: + # The traced TP x SP floor. Physical rows are padded to a multiple of + # TP; each rank saves its shard. + retained = ( + sum(-(-rows // tp) for rows, grad in group_rows if grad) + * layers + * self._hidden_size + * 2 + ) + if gradient_rows: + self._checkpoint_moe_bytes_per_token() + workspace = max( + self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) + + (0 if grad else 4 * rows * self._hidden_size * 2) + for (rows, grad), ref in zip(group_rows, refs, strict=True) + ) + if self._gdn_layers and gradient_rows: + # Recurrent states grow with segments, not rows; backward recomputes + # one layer at a time. Padding to TP adds up to TP - 1 one-token + # roots per group. Kernel-internal chunk states are not bounded here. + 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 + 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 + # and pre-MLP norm output, and its MoE stage its routing state. + mixer = ( + self._recomputed_mixer_bytes_per_token() + + ( + 2 * self._hidden_size * 2 + self._moe_checkpoint_state_bytes_per_token() + if moe + else 0 + ) + if gradient_rows + else 0 + ) workspace = max( - self._moe_workspace_bytes(rows, checkpoint_grad=grad, slot_ref=ref) - + (0 if grad else 4 * rows * self._hidden_size * 2) - for (rows, grad), ref in zip(group_rows, refs, strict=True) + self._moe_workspace_bytes( + rows, routed_rows=dispatched, checkpoint_grad=grad, slot_ref=ref + ) + + (mixer * rows if grad else 4 * rows * self._hidden_size * 2) + for (rows, grad), ref, dispatched in zip(group_rows, refs, routed, strict=True) ) - if tp > 1 and self._gdn_layers and gradient_rows: - # Recurrent states grow with segments, not rows; backward recomputes - # one layer at a time. Padding to TP adds up to TP - 1 one-token - # roots per group. Kernel-internal chunk states are not bounded here. - roots = gdn_segments + (tp - 1) * sum(grad for _, grad in group_rows) - workspace += math.ceil(roots * self._gdn_segment_layer_bytes()) + if moe: + workspace += self._te_workspace_growth_bytes() return retained, workspace +def _gradient_slots( + group_rows: Sequence[tuple[int, bool]], + slot_refs: Sequence[LoRASlotRef | None] | None, +) -> frozenset[LoRASlotRef]: + """Gradient groups' adapter slots; the base model (no name) has none.""" + return frozenset( + ref + for (_, grad), ref in zip( + group_rows, slot_refs or (None,) * len(group_rows), strict=True + ) + if grad and ref is not None and ref.name is not None + ) + + +def _pending_adapter_gradient_bytes( + self: TrainerRank, refs: Iterable[LoRASlotRef] +) -> tuple[int, ...]: + """Local adapter gradient bytes the next backward allocates, per decoder layer. + + One entry per decoder layer, then one for parameters outside the + decoder. Backward reaches a parameter's highest decoder layer first, so + a parameter shared across layers counts there once; those outside the + decoder count as live throughout. Only unallocated gradients count: + within a step, later waves find the rest in the availability baseline. + Empty when none is pending. + """ + refs = tuple(dict.fromkeys(refs)) + if not refs or len(self.runtime.model) != 1: + return () + try: + from art.megatron.lora import LoRA + except ModuleNotFoundError as error: + if error.name != "megatron": + raise + return () + chunk = self.runtime.model[0] + try: + layers = _impl._language_model(chunk).decoder.layers + except (AttributeError, RuntimeError): + return () + layer_of: dict[int, int] = {} + params: dict[int, torch.nn.Parameter] = {} + + def slot_params( + modules: Iterable[torch.nn.Module], + ) -> Iterator[torch.nn.Parameter]: + for module in modules: + # A LoRA without slot tables holds no slot parameters. + if isinstance(module, LoRA) and "_slot_keys" in vars(module): + for ref in refs: + yield from module.lora_slot_params(ref) + + for index, layer in enumerate(layers): + for param in slot_params(layer.modules()): + params[id(param)] = param + layer_of[id(param)] = max(layer_of.get(id(param), -1), index) + # Outside the decoder (a head runs its backward first) is live + # throughout, even for a parameter or module a decoder layer also uses: + # walk every path except through the layers themselves. + outside: list[torch.nn.Module] = [] + visited: set[int] = set() + pending_modules: list[torch.nn.Module] = [chunk] + while pending_modules: + module = pending_modules.pop() + if module is layers or id(module) in visited: + continue + visited.add(id(module)) + outside.append(module) + pending_modules.extend(module.children()) + for param in slot_params(outside): + params[id(param)] = param + layer_of[id(param)] = len(layers) + # A checkpoint's other trainable parameters (custom objects) have no + # decoder position; count them as live throughout. + for ref in refs: + slot = ( + None + if ref.name is None + else getattr(self, "_checkpoint_slots", {}).get(ref.name) + ) + for param in () if slot is None else slot.params: + if id(param) not in params: + params[id(param)] = param + layer_of[id(param)] = len(layers) + sizes = [0] * (len(layers) + 1) + for param_id, param in params.items(): + if ( + param.requires_grad + and param.grad is None + and getattr(param, "main_grad", None) is None + ): + sizes[layer_of[param_id]] += param.numel() * param.element_size() + return tuple(sizes) if any(sizes) else () + + +def _checkpoint_gradient_groups( + self: TrainerRank, + group_rows: Sequence[tuple[int, bool]], + slot_refs: Sequence[LoRASlotRef | None] | None, +) -> tuple[tuple[LoRASlotRef | None, tuple[int, ...]], ...]: + """Each gradient group's adapter slot and per-layer saved boundaries. + + In execution order, as ``_checkpoint_memory_floor`` prices them: every + decoder layer saves the group's rows (this rank's TP shard). The base + model (no name) has no adapter slot. + """ + tp = self._topology_key()[1] + return tuple( + ( + ref if ref is not None and ref.name is not None else None, + (-(-rows // tp) * self._hidden_size * 2,) * self._num_layers, + ) + for (rows, grad), ref in zip( + group_rows, slot_refs or (None,) * len(group_rows), strict=True + ) + if grad + ) + + +def _checkpoint_adapter_gradient_bytes( + self: TrainerRank, + groups: Sequence[tuple[LoRASlotRef | None, Sequence[int]]], + *, + head: bool = False, +) -> 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``). 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,)) + if pending and len(pending) != len(boundaries) + 1: + return 0 + chains.append((pending or (0,) * (len(boundaries) + 1), boundaries)) + if head: + return self._adapter_gradient_head(chains) + return self._adapter_gradient_walk(chains) + + +def _adapter_gradient_walk( + chains: Sequence[tuple[Sequence[int], Sequence[int]]], +) -> int: + """The adapter-gradient peak beyond the floor over gradient groups' backward. + + Each chain is a gradient group's pending gradient bytes (per decoder + layer, then outside the decoder) and saved-boundary bytes per layer. + Backward recomputes the last layer first. While it recomputes layer i it + still holds the saved boundaries of layers 0..i and every adapter + gradient allocated so far: those of layers i..L-1 (a layer allocates its + own during its backward) and any outside the decoder. The floor already + prices all L boundaries at once, so one group's extra peak is + max(0, max over i of gradients(i..) - boundaries(i+1..)), taken at every + layer over the real per-layer gradient sizes and the caller's per-layer + boundaries, not along a uniform-layer line. A short + first wave peaks at layer 0 (Qwen3.6-35B-A3B CP2: 830-900 MB of expert + LoRA gradients live at its peak), a long one at the last layer. + Groups run their backward one after another, not layer by layer + together: autograd drains the last-forwarded group's chain first, and + separate backward calls can come in either order. While one group runs, + each group already run holds all its gradients and none of its + boundaries, and each group yet to run all its boundaries. Any set of the + other groups can have run first, so the worst adds every other group + whose gradients outweigh its boundaries. + """ + nets = [sum(pending) - sum(boundaries) for pending, boundaries in chains] + others = sum(max(0, net) for net in nets) + worst = 0 + for (pending, boundaries), net in zip(chains, nets, strict=True): + extra = gradients = pending[-1] + released = 0 + for index in range(len(boundaries) - 1, -1, -1): + gradients += pending[index] + extra = max(extra, gradients - released) + released += boundaries[index] + worst = max(worst, extra + others - max(0, net)) + return worst + + +def _triton_min_rows(self: TrainerRank) -> int: + """Rows the fused head statistics need in a chunk (process-wide setting).""" + return int(_impl.os.environ.get("ART_TRAINER_RANK_TRITON_MIN_ROWS", "64")) + + +def _head_backward_traced( + self: TrainerRank, + requests: Sequence[AnyForwardInput], + rows: int, + upper_rows: int | None = None, +) -> bool | None: + """Whether a gradient group's head backward is the traced one. + + The head stage (``_checkpoint_head_stage_bytes``) relies on + ``_group_head_workspace_bytes`` bounding the head's backward buffers. + Qwen3.6-35B-A3B CP2 traces establish that for target-only requests on + the standard logit scale through the fused Triton statistics, which + need at least ``ART_TRAINER_RANK_TRITON_MIN_ROWS`` rows in the first + projected chunk. Top-k, logits and hidden-state outputs keep further + dense gradients, and the FP32 fallback wider copies; other CP sizes + are untraced. The fused statistics must have run in this process and + never fallen back after an error; a plan priced on them then runs them + strictly (``_plan_head_backward_traced``), so a later failure raises + rather than taking the wider FP32 path. ``rows`` bounds the first chunk's rows from below, + ``upper_rows`` from above: None when they straddle the threshold. + A CP rank projecting fewer rows falls back on its own; the head stage + prices that (``_HEAD_FALLBACK_BUFFERS``). ``ART_TRAINER_RANK_TRITON_*`` + settings are process-wide: changing them while a plan is in flight is + unsupported (a staged chunk could then fall back unpriced). + """ + if ( + not requests + or any( + request.target_tokens is None + # Several labels per row save gather indices and masks per label; + # one label per input token, in either accepted layout, does not. + or request.target_tokens.numel() != request.input_tokens.numel() + or request.top_k is not None + or request.logits + or request.hidden_states + for request in requests + ) + or _impl.os.environ.get("ART_TRAINER_RANK_TRITON_TOPK", "1").lower() + in {"0", "false"} + # Target-only heads run the logsumexp kernel; another kernel's + # success does not prove it. + or "local_logsumexp_stats" not in _impl._TRITON_STATS_STATE["succeeded"] + or _impl._TRITON_STATS_STATE["failed"] + or self._topology_key()[2] != 2 + or not self._head_workspace_bytes(1) + or not _head_target_backward(self) + ): + return False + minimum = self._triton_min_rows() + if rows >= minimum: + return True + if upper_rows is None or upper_rows < minimum: + return False + return None + + +def _te_workspace_growth_bytes(self: TrainerRank) -> int: + """Transformer Engine's cuBLAS workspaces, until its GEMMs allocate them. + + The first plain and grouped GEMMs allocate them during a call, and TE + keeps them for the process; later calls see them as used memory. + """ + try: + from transformer_engine.pytorch.cpp_extensions import gemm + except ImportError: + return _impl._TE_CUBLAS_WORKSPACE_BYTES + info = getattr(gemm.get_cublas_workspace, "cache_info", None) + if callable(info) and info().currsize >= 2: + return 0 + return _impl._TE_CUBLAS_WORKSPACE_BYTES + + +def _moe_checkpoint_state_bytes_per_token(self: TrainerRank) -> int: + """Per local token beside the recomputed MoE stage's routed rows. + + FP32 router scores and the boolean routing map; the dispatcher's state + (the EP1 permutation's int32 row-id map of 2E + 1, or HybridEP's FP32 + probability copy and handle metadata of about 5E); and the shared + expert's saved FC1 gate/up and GLU outputs. Qwen3.6-35B-A3B traces: + 3,332 and 3,593 bytes of routing state at EP1 and EP2, 3 KB shared. + """ + geometry = self._geometry + experts = geometry.moe_experts + if not experts: + return 0 + routing = experts * (4 + 1) + ( + 4 * (2 * experts + 1) if self._parallel_shape.ep == 1 else 9 * experts + 16 + ) + return routing + 3 * geometry.moe_shared_expert_ffn * self._param_dtype_size + + +def _moe_recompute_covered_for( + self: TrainerRank, slot_ref: "LoRASlotRef | None" +) -> bool: + """Whether an explicit slot keeps the constructor's full MoE coverage.""" + if not getattr(self, "_moe_recompute_covered", False): + return False + if slot_ref is None or slot_ref.name is None: + return True + enclosed: list[bool] = [] + coefficient = _impl._moe_output_bytes_per_token( + self.runtime.model, + self._parallel_shape, + checkpoint_grad=True, + converted_stages=[], + slot_ref=slot_ref, + enclosed=enclosed, + ) + # The constructor's walk already matched every decoder layer. + return ( + coefficient > 0 + and len(enclosed) == len(getattr(self, "_moe_gradient_enclosed", ())) + and all(enclosed) + ) + + +def _checkpoint_input_gradient_bytes( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None = None, + *, + retained: int | None = None, +) -> int: + """Gradient rows live at the recomputed layer's peak. + + Backward recomputes the last layer first, so its peak meets every saved + boundary but only the one incoming gradient. Where the MoE stage covers + every layer's recompute, FC1 included (Qwen3.6-35B-A3B traces at CP1, + CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere, including above + CP2 (more remote attention stages than the mixer's CP2 allowance), keep + one gradient per boundary: that allowance also covers dense MLP and + other recompute work the floor does not price. ``retained`` is the + floor's boundary charge where the caller already has it. + """ + if retained is None: + retained, _ = self._checkpoint_memory_floor(group_rows) + if not retained: + return 0 + if self._checkpoint_gradient_covered(group_rows, slot_refs): + return sum(rows for rows, grad in group_rows if grad) * self._hidden_size * 2 + return retained + + +def _checkpoint_gradient_covered( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, +) -> bool: + """Whether the traced MoE stage covers every gradient group's recompute.""" + refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + return bool( + self._topology_key()[2] <= 2 + and self._checkpoint_moe_bytes_per_token() + and all( + self._moe_recompute_covered_for(ref) + for (_, grad), ref in zip(group_rows, refs, strict=True) + if grad + ) + ) + + +def _checkpoint_head_stage_bytes( + self: TrainerRank, + head_workspace_bytes: int, + gradient: int, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None, +) -> int | None: + """The head backward's peak beyond the boundaries and ``gradient``. + + The head finishes its backward before the decoder's recompute starts, + so its buffers never meet a recomputed layer's workspace or the adapter + gradients that recompute allocates. Where the MoE stage covers the + recompute, ``gradient`` is the gradient groups' rows x 2H, and each of + these is at most that: the final decoder outputs, the hidden rows each + checkpointed head chunk saved (this rank's rows), and the hidden-row + gradient. Beside them are each row's RoPE embedding and index state, + the head's own buffers, TE's cuBLAS workspaces from the forward's first + GEMMs, and the adapter gradients of groups whose backward ran first. + Qwen3.6-35B-A3B CP2 traces (EP1 and EP2, single and multi-request + waves) show these terms at the head's peak. Elsewhere None: the head + 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``. + """ + 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) + # 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( + min( + self._triton_min_rows() - 1, + _impl._HEAD_CHUNK_TOKENS, + ) + ) + return ( + max(head_workspace_bytes, fallback) + + 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 + ) + ) + + +def _backward_row_state_bytes(self: TrainerRank) -> int: + """Per local row, what the backward keeps beside hidden-width tensors. + + The FP32 RoPE embedding at the rotary width (256 bytes for + Qwen3.6-35B-A3B) and ``_BACKWARD_ROW_STATE_BYTES`` of index state. + """ + rope = getattr(_impl._language_model(self.runtime.model[0]), "rotary_pos_emb", None) + frequencies = getattr(rope, "inv_freq", None) + width = ( + 2 * frequencies.numel() if isinstance(frequencies, _impl.torch.Tensor) else 0 + ) + return width * 4 + _impl._BACKWARD_ROW_STATE_BYTES + + +def _mixer_activation_widths(self: TrainerRank) -> tuple[float, float]: + """Saved activations per token of one attention and one GDN layer. + + Elements, not bytes: what a layer's attention or GDN mixer keeps for + its backward, including the input norm output (and gathered LoRA + inputs under sequence parallelism). + """ + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + sp = tp if self._sequence_parallel else 1 + # Gathered LoRA inputs alias norm output without sequence sharding. + gathered = hidden if sp > 1 else 0 + common = 2 * hidden / sp + gathered + attention_width = geometry.num_attention_heads * geometry.kv_channels or hidden + kv_width = geometry.num_query_groups * geometry.kv_channels or hidden + gated = self._attention_output_gate + attention = common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp + if 0 < geometry.num_query_groups < tp: + # SelfAttentionLinearQKVLoRA constructs global QKV before + # slicing it when KV groups cannot be partitioned across TP. + attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( + 1 - 1 / tp + ) + gdn = ( + common + + ( + 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim + + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim + ) + / tp + ) + return attention, gdn + + +def _recomputed_mixer_bytes_per_token(self: TrainerRank) -> int: + """Saved mixer activations of the layer being recomputed, per row. + + Full recompute replays one layer with gradients, and its attention or + GDN mixer keeps what its backward needs across that layer's MoE stage. + Price the larger mixer the model has. Sites and sizes come from + Qwen3.6-35B-A3B allocator traces on H200, per local token at the + layer's recompute peak: + + - attention: the retained attention width above (66 KB measured at + CP1, 69 KB priced). A context-parallel rank also keeps its + stage-padded Q/K/V, the stage output and a core-attention copy + (94 KB measured at CP2, 95 KB priced). CP above 2 uses the CP2 + allowance; ranks with several remote stages may keep more. + - GDN: norm output, the projected q/k/v, their l2norm outputs + expanded to the value heads, five more value-width tensors (z, two + segment-layout tensors, the gated-norm output and its gated + product) and the chunk decay matrix: 80 KB measured at CP1, 82 KB + priced. A context-parallel rank's exchanged layout holds about 7% + more rows; its hidden-width input exchange and value-width output + allowance price that (88 KB measured at CP2, 94 KB priced). + """ + geometry = self._geometry + hidden = self._hidden_size + tp = max(1, self._topology_key()[1]) + cp = self._topology_key()[2] > 1 + widths = [] + if self._gdn_layers < self._num_layers: + attention, _gdn = self._mixer_activation_widths() + if cp: + 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) + if self._gdn_layers: + key = geometry.gdn_key_heads * geometry.gdn_key_head_dim + value = geometry.gdn_value_heads * geometry.gdn_value_head_dim + normalized = 2 * geometry.gdn_value_heads * geometry.gdn_key_head_dim + chunk = 64 * geometry.gdn_value_heads + 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) + + +def _adapter_gradient_head( + chains: Sequence[tuple[Sequence[int], Sequence[int]]], +) -> int: + """Adapter gradients live while one group's head runs its backward. + + Chains as ``_adapter_gradient_walk``. None of the group's decoder layers + has run yet; its gradients outside the decoder are live, and so is every + other group whose gradients outweigh its boundaries, as any may have run + first. + """ + nets = [max(0, sum(pending) - sum(boundaries)) for pending, boundaries in chains] + others = sum(nets) + return max( + ( + pending[-1] + others - net + for (pending, _), net in zip(chains, nets, strict=True) + ), + default=0, + ) + + def _retained_memory_bytes( self: TrainerRank, signature: _impl._MemorySignature, @@ -674,6 +1285,7 @@ def _estimate_flat_forward( memory_minimal: bool = False, sync_planning_errors: bool = False, gdn_segments: list[int] | None = None, + head_traced: list[bool | None] | None = None, ) -> tuple[int, int, _impl._MemorySignature, tuple[tuple[int, bool], ...], int] | None: """Estimate packed tokens for width probing. @@ -689,6 +1301,8 @@ def _estimate_flat_forward( ``gdn_segments`` receives each gradient group's segment count: exact layouts' actual counts; in cheap mode, the same kind of bound as the token count (twice the requests, as a radix tree has fewer, or one). + ``head_traced`` receives each gradient group's ``_head_backward_traced`` + over its projected-row bounds. """ if sync_planning_errors: @@ -705,6 +1319,19 @@ def _estimate_flat_forward( # This cheap return type has no slot metadata. Materialize the # exact plan instead of admitting with the constructor rank. return None + gradient_slots = [ + ref + for (ref, grad), _ in groups + if grad and ref is not None and ref.name is not None + ] + if ( + gradient_slots + and getattr(self, "_recompute_granularity", None) == "full" + and self._pending_adapter_gradient_bytes(gradient_slots) + ): + # The step's first backward allocates adapter gradients that + # only slot metadata can price; the exact plan carries it. + return None if ( any(mode for (_, mode), _ in groups) and _impl._gdn_memory.model_shapes(self) is not None @@ -724,6 +1351,10 @@ def _estimate_flat_forward( head_requests = tuple(requests[index] for index in group_indices) lower = self._head_projection_rows(head_requests, lower_bound=True) upper = self._head_projection_rows(head_requests) + if grad_enabled and head_traced is not None: + head_traced.append( + self._head_backward_traced(head_requests, lower, upper) + ) if exact: tree, layout = self._select_group_layout( tuple( @@ -952,8 +1583,10 @@ def _memory_check( logical_tokens=forward.active_logical_tokens, gdn_segments=forward.grad_segment_count, group_rows=self._plan_group_rows(forward), + group_routed_rows=self._plan_group_routed_rows(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), checkpoint_floor=_impl._gdn_memory.plan_floor(self, forward), retained_tokens=self._plan_retained_tokens(forward), ) + int(self._plan_hybridep_growth_bytes(forward) * _impl._MEMORY_SAFETY_FACTOR) @@ -1064,8 +1697,10 @@ def _estimate_required_memory_bytes_from_values( logical_tokens: int | None = None, gdn_segments: int = 0, group_rows: tuple[tuple[int, bool], ...] = (), + group_routed_rows: tuple[int, ...] | None = None, slot_refs: tuple["LoRASlotRef | None", ...] | None = None, head_workspace_bytes: int = 0, + head_backward_traced: bool = False, checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, @@ -1086,24 +1721,7 @@ def _estimate_required_memory_bytes_from_values( # Gathered LoRA inputs alias norm output without sequence sharding. gathered = hidden if sp > 1 else 0 common = 2 * hidden / sp + gathered - attention_width = geometry.num_attention_heads * geometry.kv_channels or hidden - kv_width = geometry.num_query_groups * geometry.kv_channels or hidden - gated = self._attention_output_gate - attention = common + ((7 if gated else 5) * attention_width + 3 * kv_width) / tp - if 0 < geometry.num_query_groups < tp: - # SelfAttentionLinearQKVLoRA constructs global QKV before - # slicing it when KV groups cannot be partitioned across TP. - attention += ((2 if gated else 1) * attention_width + 2 * kv_width) * ( - 1 - 1 / tp - ) - gdn = ( - common - + ( - 4 * geometry.gdn_key_heads * geometry.gdn_key_head_dim - + 8 * geometry.gdn_value_heads * geometry.gdn_value_head_dim - ) - / tp - ) + attention, gdn = self._mixer_activation_widths() ffn_width = geometry.ffn_hidden_size or 4 * hidden mlp = common + self._mlp_activation_factor * ffn_width / tp if geometry.moe_experts: @@ -1152,29 +1770,48 @@ def _estimate_required_memory_bytes_from_values( ) # Groups execute sequentially: summed packed rows conservatively bound # this FC2 component, not all workspace or retained graphs. + grouped = signature.topology[2] > 1 and bool(group_rows) static_compute = max( static_compute, *( self._moe_workspace_bytes( - sum(rows for rows, _ in group_rows) - if signature.topology[2] > 1 and group_rows - else packed_tokens, + sum(rows for rows, _ in group_rows) if grouped else packed_tokens, + routed_rows=sum(group_routed_rows) + if grouped and group_routed_rows is not None + else None, slot_ref=ref, ) for ref in (slot_refs or (None,)) ), ) retained, workspace = ( - self._checkpoint_memory_floor(group_rows, slot_refs, gdn_segments) + self._checkpoint_memory_floor( + group_rows, slot_refs, gdn_segments, routed_rows=group_routed_rows + ) if checkpoint_memory is None else checkpoint_memory ) - static_compute = max( - static_compute, - max(retained, checkpoint_floor[0]) - + max(workspace, head_workspace_bytes, checkpoint_floor[1]) - + (retained if include_checkpoint_input_gradient else 0), - ) + decoder_workspace = max(workspace, checkpoint_floor[1]) + peak = max(decoder_workspace, head_workspace_bytes) + if include_checkpoint_input_gradient and retained: + # 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) + ) + head_stage = self._checkpoint_head_stage_bytes( + head_workspace_bytes, gradient, group_rows, slot_refs + ) + peak += adapter_gradient + if head_stage is not None: + # An untraced head keeps the unstaged price as a floor. + staged = max(decoder_workspace + adapter_gradient, head_stage) + peak = staged if head_backward_traced else max(peak, staged) + peak += gradient + if profiled is None and any(grad for _, grad in group_rows): + peak += _impl._COLD_RECOMPUTE_TRANSIENT_BYTES + static_compute = max(static_compute, max(retained, checkpoint_floor[0]) + peak) if signature.topology[2] > 1: # Local head results coexist with full CP outputs during gathering. # Uneven rank plans can assign all of an item's rows to one rank. diff --git a/src/art/trainer_rank/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index 2a077b57d..5e2b7bfa8 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -431,6 +431,7 @@ def _split_chunk_lower_cost( packed_tokens = 0 unshared_packed_tokens = 0 head_workspace_bytes = 0 + head_traced: list[bool | None] = [] group_rows: list[tuple[int, bool]] = [] for (_slot, grad_enabled), group_indices in groups: estimated = _impl.estimate_prefix_tree_packed_tokens( @@ -444,15 +445,24 @@ def _split_chunk_lower_cost( cp = max(1, self._topology_key()[2]) group_rows.append((-(-physical_rows // cp), grad_enabled)) head_requests = tuple(requests[index] for index in group_indices) + lower = self._head_projection_rows(head_requests, lower_bound=True) head_workspace_bytes = max( head_workspace_bytes, self._group_head_workspace_bytes( - self._head_projection_rows(head_requests, lower_bound=True), + lower, head_requests, grad_enabled=grad_enabled, lower_bound=True, ), ) + if grad_enabled: + head_traced.append( + self._head_backward_traced( + head_requests, + lower, + self._head_projection_rows(head_requests), + ) + ) unshared_packed_tokens += self._physical_tokens( sum(int(rows[index].numel()) for index in group_indices) ) @@ -464,17 +474,27 @@ def _split_chunk_lower_cost( slot_groups=tuple(key for key, _ in groups), ) logical_tokens = _impl._active_logical_tokens(requests) - cost = self._subforward_cost( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - group_rows=tuple(group_rows), - 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. - retained_tokens=(packed_tokens + signature.topology[2] - 1) - // signature.topology[2], + # 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( + ( + self._subforward_cost( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + group_rows=tuple(group_rows), + slot_refs=tuple(ref for (ref, _), _ in groups), + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + # The average CP load is an optimistic bound, not an + # admission cost. + retained_tokens=(packed_tokens + signature.topology[2] - 1) + // signature.topology[2], + ) + for traced in _impl._traced_states(head_traced) + ), + key=lambda cost: cost.required, ) profile = self._memory_profiles.get(signature) if ( @@ -535,6 +555,76 @@ def _plan_group_rows( ) +def _plan_head_backward_traced(self: TrainerRank, plan: _FlatForwardPlan) -> bool: + """Every gradient group's head backward is traced (``_head_backward_traced``). + + When that stages the plan's head price (``_checkpoint_head_stage_bytes`` + applies), marks the plan: its gradient groups then run the fused + statistics strictly (``_execute_flat_plan``). The mark stays once set: + a later re-pricing cannot weaken the admission that relied on it. + Every plan admitted on a staged price has been priced here: staging + needs CP2, where the cheap width estimate declines and admission + prices the materialized plan (``_estimate_flat_forward``). Strictness + trades the FP32 fallback's availability for memory safety: a fused + kernel error then fails the wave, and a CP peer waits in its next + collective like after a rank-local OOM. + """ + traced = [ + self._head_backward_traced( + requests, + self._head_projection_rows( + requests, + positions=group.packed.positions_by_sequence, + lower_bound=True, + ), + ) + for group in plan.groups + if group.grad_enabled + for requests in (tuple(item.request for item in group.items),) + ] + eligible = bool(traced) and all(state is True for state in traced) + if ( + eligible + and self._plan_head_workspace_bytes(plan) + and self._checkpoint_gradient_covered( + self._plan_group_rows(plan), tuple(g.slot_ref for g in plan.groups) + ) + ): + object.__setattr__(plan, "_head_staged", True) + return eligible + + +def _plan_group_routed_rows( + self: TrainerRank, plan: _FlatForwardPlan +) -> tuple[int, ...]: + """Rows one rank's experts receive per group at balanced routing. + + HybridEP dispatches the whole EP group's rows. When that group is this + rank's CP group, a balanced rank receives the group's rows over EP, + however unevenly CP split them; otherwise keep the local rows. + """ + rows = self._plan_group_rows(plan) + if ( + not getattr(self, "_ep_group_is_cp_group", False) + or plan.signature.topology[2] <= 1 + ): + return tuple(local for local, _ in rows) + topology = self._topology() + return tuple( + min( + local, + -( + -self._cp_group_model_tokens( + _impl._pad_packed_batch(group.packed, multiple=int(topology.tp)), + topology=topology, + ) + // int(topology.cp) + ), + ) + for (local, _), group in zip(rows, plan.groups, strict=True) + ) + + def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -543,8 +633,10 @@ def _plan_cost(self: TrainerRank, plan: _FlatForwardPlan) -> _SubforwardCost: logical_tokens=plan.active_logical_tokens, gdn_segments=plan.grad_segment_count, group_rows=self._plan_group_rows(plan), + group_routed_rows=self._plan_group_routed_rows(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), checkpoint_floor=_impl._gdn_memory.plan_floor(self, plan), retained_tokens=self._plan_retained_tokens(plan), hybridep_growth_bytes=self._plan_hybridep_growth_bytes(plan), @@ -679,11 +771,13 @@ def estimate(width: int) -> tuple[_MemoryCheck, bool, bool] | None: indices, local_inputs = local_slice(width) local_requests = list(_impl._flatten(local_inputs)) cheap_segments: list[int] = [] + cheap_traced: list[bool | None] = [] values = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, sync_planning_errors=True, gdn_segments=cheap_segments, + head_traced=cheap_traced, ) if not self._all_ranks_true(values is not None): estimates[width] = None @@ -699,18 +793,27 @@ def priced( head_workspace_bytes: int, *, gdn_segments: int, + head_traced: Sequence[bool | None], + lower: bool, ) -> tuple[_MemoryCheck, int, int, _MemorySignature]: with self._planning_status(True): - required = self._estimate_required_memory_bytes_from_values( - packed_tokens=packed_tokens, - output_bytes=output_bytes, - signature=signature, - logical_tokens=logical_tokens, - # Gradient groups' segments: exact layouts' counts, else - # a bound matching the estimate's (_estimate_flat_forward). - gdn_segments=gdn_segments, - group_rows=group_rows, - head_workspace_bytes=head_workspace_bytes, + # A bound over layouts whose head may or may not stage: + # the higher price to accept, the lower to reject. + required = (min if lower else max)( + self._estimate_required_memory_bytes_from_values( + packed_tokens=packed_tokens, + output_bytes=output_bytes, + signature=signature, + logical_tokens=logical_tokens, + # Gradient groups' segments: exact layouts' counts, + # else a bound matching the estimate's + # (_estimate_flat_forward). + gdn_segments=gdn_segments, + group_rows=group_rows, + head_workspace_bytes=head_workspace_bytes, + head_backward_traced=traced, + ) + for traced in _impl._traced_states(head_traced) ) return ( self._memory_check_required(required, sync_across_dp=True), @@ -723,6 +826,7 @@ def priced_estimate( *, exact: bool, memory_minimal: bool ) -> tuple[_MemoryCheck, int, int, _MemorySignature] | None: segments: list[int] = [] + traced: list[bool | None] = [] estimated = self._estimate_flat_forward( local_requests, checkpoint=checkpoint, @@ -730,11 +834,18 @@ def priced_estimate( memory_minimal=memory_minimal, sync_planning_errors=True, gdn_segments=segments, + head_traced=traced, ) return ( None if estimated is None - else priced(*estimated, gdn_segments=sum(segments)) + else priced( + *estimated, + gdn_segments=sum(segments), + head_traced=traced, + # Only the cheap full-sharing count rejects. + lower=memory_minimal and not exact, + ) ) def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: @@ -748,7 +859,12 @@ def trusted(packed_tokens: int, signature: _MemorySignature) -> bool: # reject on memory, or when it would reject on profile trust while # a profile exists — the selected layout may be far smaller than # the bound and squarely inside the profiled regime. - selected = priced(*values, gdn_segments=sum(cheap_segments)) + selected = priced( + *values, + gdn_segments=sum(cheap_segments), + head_traced=cheap_traced, + lower=False, + ) profiled = self._all_ranks_true(selected[3] in self._memory_profiles) needs_exact = not selected[0].fits or ( profiled and not trusted(selected[1], selected[3]) @@ -1415,6 +1531,8 @@ def _fill_planner_snapshot( "gdn_segments": child.grad_segment_count, "retained_tokens": self._plan_retained_tokens(child), "group_rows": self._plan_group_rows(child), + "group_routed_rows": self._plan_group_routed_rows(child), + "head_backward_traced": self._plan_head_backward_traced(child), "hybridep_growth_bytes": ( self._plan_hybridep_growth_bytes(child) ), diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index bb7e8297c..f581ae9b4 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -15,7 +15,7 @@ from collections.abc import Callable, Mapping from contextlib import nullcontext -from dataclasses import asdict, dataclass +from dataclasses import MISSING, asdict, dataclass from datetime import datetime, timezone import hashlib from itertools import islice @@ -600,6 +600,18 @@ def replay( raise ValueError( "incomplete replay: immutable rank fields differ (including MoE stages)" ) + layers = values["num_layers"] + if type(layers) is not int or not 0 < layers <= _planner_replay.MAX_LAYERS: + raise ValueError("incomplete replay: recorded layer count out of bounds") + # Recomputed-layer pricing reads these alongside the geometry below. + if any( + type(values[name]) is not int or not 0 <= values[name] < 2**63 + for name in ("hidden_size", "param_dtype_size", "gdn_layers") + ) or any( + type(values[name]) is not bool + for name in ("sequence_parallel", "attention_output_gate") + ): + raise ValueError("incomplete replay: recorded rank dimensions are invalid") rank = _planner_replay.ReplayRank.__new__(_planner_replay.ReplayRank) for name in _RANK_FIELDS - {"one_layer_recompute"}: setattr(rank, "_" + name, values[name]) @@ -607,7 +619,19 @@ def replay( raise ValueError("incomplete replay: recompute mode is not recorded") rank._recorded_one_layer_recompute = values["one_layer_recompute"] rank._moe_forward_stages = tuple(tuple(row) for row in values["moe_forward_stages"]) - rank._geometry = ModelGeometry(**values["geometry"]) + # Recomputed-layer pricing multiplies these; accept only what live + # construction (ModelGeometry.from_config) records. Reports may omit + # fields that default to zero. + geometry = values["geometry"] + fields = ModelGeometry.__dataclass_fields__ + required = {name for name, field in fields.items() if field.default is MISSING} + if ( + type(geometry) is not dict + or not required <= set(geometry) <= set(fields) + or any(type(v) is not int or not 0 <= v < 2**63 for v in geometry.values()) + ): + raise ValueError("incomplete replay: recorded model geometry is invalid") + rank._geometry = ModelGeometry(**geometry) dp, tp, cp, pp = values["topology"] rank._topology_key = lambda: (dp, tp, cp, pp) # Check the same shared subforward inventories as capture, including the diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index f0da8bbe1..2c9d9e747 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -11,16 +11,29 @@ from dataclasses import asdict import json from types import MethodType, SimpleNamespace -from typing import Any +from typing import Any, NamedTuple from . import _gdn_memory, _impl, _memory _MAX_BYTES = 262144 _MAX_GROUPS = 1024 +# Replay sizes per-layer tuples from this recorded rank field; bound it as capture does. +MAX_LAYERS = 1024 _MAX_SEGMENTS = 4096 _MAX_INPUT_VALUES = 1_000_000 +# Estimators bound on TrainerRank as plain functions, not methods. +_STATIC_ESTIMATORS = frozenset( + { + "_split_required_memory", + "_gradient_slots", + "_adapter_gradient_walk", + "_adapter_gradient_head", + } +) + + _REFUSALS = frozenset( { "runtime_group_inventory_over_limit", @@ -41,6 +54,15 @@ ) +# Frozen per plan; None where no checkpointed decoder would read them. +_RECOMPUTE_READERS = ( + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", +) + + def refusal_reason(error: Exception) -> str: # Never retain arbitrary reader exception text, model reprs or token values. if ( @@ -68,7 +90,10 @@ def reserve(size: int) -> None: def capture(rank: Any, plan: Any) -> dict[str, Any]: - if rank._num_layers > 1024 or not 0 < len(plan.groups) <= _MAX_GROUPS: + if ( + not 0 < rank._num_layers <= MAX_LAYERS + or not 0 < len(plan.groups) <= _MAX_GROUPS + ): raise ValueError("runtime_group_inventory_over_limit") if ( getattr(rank.runtime.provider, "expert_model_parallel_size", 1) > 1 @@ -96,12 +121,32 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_physical_tokens", "_plan_group_rows", "_plan_retained_tokens", + "_gradient_slots", + "_pending_adapter_gradient_bytes", + "_checkpoint_gradient_groups", + "_checkpoint_adapter_gradient_bytes", + "_adapter_gradient_walk", + "_adapter_gradient_head", + "_checkpoint_input_gradient_bytes", + "_checkpoint_gradient_covered", + "_moe_recompute_covered_for", + "_checkpoint_head_stage_bytes", + "_backward_row_state_bytes", + "_te_workspace_growth_bytes", + "_moe_checkpoint_state_bytes_per_token", + "_recomputed_mixer_bytes_per_token", + "_mixer_activation_widths", + "_triton_min_rows", + # Producers of recorded arguments replay checks against the facts. + "_plan_group_routed_rows", + "_plan_head_backward_traced", + "_head_backward_traced", ): method = getattr(rank, name) expected = getattr(_impl.TrainerRank, name) supported = ( method is expected - if name == "_split_required_memory" # The sole static estimator. + if name in _STATIC_ESTIMATORS else type(method) is MethodType and method.__self__ is rank and method.__func__ is expected @@ -131,19 +176,19 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: head_values = _MAX_INPUT_VALUES def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: - coefficient, stages = _memory._moe_workspace_terms( + coefficient, stages, shared = _memory._moe_workspace_terms( rank, checkpoint_grad=checkpoint_grad, slot_ref=ref ) if len(stages) > 4096: raise ValueError("runtime_stage_inventory_over_limit") - reserve(64 + 72 * len(stages)) + reserve(96 + 72 * len(stages)) if ( type(coefficient) is not int or not 0 <= coefficient < 2**63 or any(v >= 2**63 for stage in stages for v in stage) ): raise ValueError("runtime_dimension_unsupported") - return [coefficient, [list(stage) for stage in stages]] + 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( @@ -214,6 +259,18 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: if backwards else 0 ) + adapter = None + if group.grad_enabled and name is not None: + # The live estimator reads this slot's unallocated gradient bytes + # per decoder layer; freeze them with the selection. + kind = getattr(group.slot_ref, "kind", None) + if kind is not None and (type(kind) is not str or len(kind) > 64): + raise ValueError("runtime_slot_identity_unsupported") + pending = rank._pending_adapter_gradient_bytes((group.slot_ref,)) + if len(pending) > 1025: + 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]} model = _gdn_memory.model_shapes(rank, group.slot_ref) if has_grad else None if model is not None: if len(model[1]) > 1024: @@ -231,6 +288,9 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: "gradient": terms(True, group.slot_ref), "head_rows": projected, "head_target_rows": target_rows, + "adapter": adapter, + "moe_covered": group.grad_enabled + and rank._moe_recompute_covered_for(group.slot_ref), "gdn": None if model is None else { @@ -252,12 +312,17 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: }, } ) + layers = _memory._checkpoint_layers(rank, rank._plan_group_rows(plan)) facts = { - "version": 1, - "checkpoint_layers": _memory._checkpoint_layers( - rank, rank._plan_group_rows(plan) - ), + "version": 3, + "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, + # read (like live pricing) only with a checkpointed decoder. + **{ + name: getattr(rank, "_" + name)() if layers else None + for name in _RECOMPUTE_READERS + }, "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -285,12 +350,16 @@ def integer(value: Any, *, minimum: int = 0) -> None: "version", "checkpoint_layers", "checkpoint_moe_bytes_per_token", + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", "head_vocabulary", "head_target_backward", "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 1: + if type(facts["version"]) is not int or facts["version"] != 3: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", @@ -298,6 +367,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: "head_vocabulary", ): integer(facts[key]) + for key in _RECOMPUTE_READERS: + if not facts["checkpoint_layers"]: + if facts[key] is not None: + raise ValueError("invalid runtime dimension") + continue + # Any integer Triton setting is live-accepted (a non-positive one + # prices no fallback chunk rows). + integer(facts[key], minimum=-(2**63) if key == "triton_min_rows" else 0) if type(facts["head_target_backward"]) is not bool: raise ValueError("invalid head backward eligibility") groups = facts["groups"] @@ -319,6 +396,8 @@ def integer(value: Any, *, minimum: int = 0) -> None: "gradient", "head_rows", "head_target_rows", + "adapter", + "moe_covered", "gdn", }, ) @@ -330,6 +409,11 @@ def integer(value: Any, *, minimum: int = 0) -> None: or len(group["slot"]) > 4096 ): raise ValueError("invalid runtime group identity") + # Only gradient groups recompute; the live reader is not asked otherwise. + if type(group["moe_covered"]) is not bool or ( + group["moe_covered"] and not group["grad"] + ): + raise ValueError("invalid MoE recompute coverage") if ( type(group["request_indices"]) is not list or len(group["request_indices"]) > 4096 @@ -351,18 +435,41 @@ def integer(value: Any, *, minimum: int = 0) -> None: terms = group[key] if ( type(terms) is not list - or len(terms) != 2 + or len(terms) != 3 or type(terms[1]) is not list or len(terms[1]) > 4096 ): raise ValueError("invalid MoE terms") - reserve(64 + 72 * len(terms[1])) + reserve(96 + 72 * len(terms[1])) integer(terms[0]) for stage in terms[1]: if type(stage) is not list or len(stage) != 2: raise ValueError("invalid MoE stage") for value in stage: integer(value) + # The shared expert's part of the coefficient (live invariant). + integer(terms[2]) + if terms[2] > terms[0]: + raise ValueError("invalid MoE terms") + adapter = group["adapter"] + if adapter is not None: + fields(adapter, {"kind", "name", "pending"}) + if ( + not group["grad"] + or (adapter["kind"] is not None and type(adapter["kind"]) is not str) + or len(adapter["kind"] or "") > 64 + or type(adapter["name"]) is not str + or len(adapter["name"]) > 4096 + or type(adapter["pending"]) is not list + or len(adapter["pending"]) > 1025 + ): + raise ValueError("invalid adapter gradient facts") + reserve(128 + 12 * len(adapter["name"]) + 24 * len(adapter["pending"])) + for value in adapter["pending"]: + integer(value) + # A slot without a kind (megatron-less reference) has none pending. + if adapter["kind"] is None and any(adapter["pending"]): + raise ValueError("invalid adapter gradient facts") gdn = group["gdn"] if gdn is not None: fields(gdn, {"layers", "shapes", "segments"}) @@ -390,6 +497,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: ) for value in segment.values(): integer(value) + # Live slot references all have a kind, or (without megatron) none do. + kinds = { + group["adapter"]["kind"] is None + for group in groups + if group["adapter"] is not None + } + if len(kinds) > 1: + raise ValueError("invalid adapter gradient facts") if len(json.dumps(facts, separators=(",", ":"))) > _MAX_BYTES: raise ValueError("runtime_facts_over_limit") @@ -420,14 +535,58 @@ def count(value: Any, depth: int = 0) -> None: count(value) +class _ReplaySlot(NamedTuple): + """A frozen adapter slot identity (the live LoRASlotRef's kind and name).""" + + kind: str | None + name: str + + class ReplayRank(_impl.TrainerRank): """The real estimator with runtime metadata readers replaced by frozen facts.""" _facts: dict[str, Any] | None = None + def _replay_slots(self, slot_refs: Any) -> Any: + # Replay passes each group's index; map it to that group's frozen slot. + if self._facts is None or slot_refs is None: + return slot_refs + groups = self._facts["groups"] + return tuple( + None + if (adapter := groups[index]["adapter"]) is None + else _ReplaySlot(adapter["kind"], adapter["name"]) + for index in slot_refs + ) + + def _gradient_slots(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._gradient_slots(group_rows, self._replay_slots(slot_refs)) + + def _checkpoint_gradient_groups(self, group_rows: Any, slot_refs: Any) -> Any: + return _memory._checkpoint_gradient_groups( + self, group_rows, self._replay_slots(slot_refs) + ) + + def _pending_adapter_gradient_bytes(self, refs: Any) -> tuple[int, ...]: + if self._facts is None: + return _memory._pending_adapter_gradient_bytes(self, refs) + refs = tuple(dict.fromkeys(refs)) + if not refs: + return () + if len(refs) != 1: + raise ValueError("replayed adapter gradients are frozen per slot") + for group in self._facts["groups"]: + adapter = group["adapter"] + if adapter is not None and (adapter["kind"], adapter["name"]) == tuple( + refs[0] + ): + return tuple(adapter["pending"]) + return () + def _head_workspace_bytes(self, rows: int) -> int: assert self._facts is not None - return _memory._dense_head_bytes(self._facts["head_vocabulary"], rows) + vocabulary = self._facts["head_vocabulary"] + return _memory._dense_head_bytes(vocabulary, rows) if rows > 0 else 0 def verify_group( self, group: dict[str, Any], layout: Any, records: list[dict[str, Any]] @@ -492,29 +651,72 @@ def verify_group( raise ValueError("head row facts disagree with selected requests/layout") def _moe_workspace_bytes( - self, rows: int, *, checkpoint_grad: bool = False, slot_ref: Any = None + self, + rows: int, + *, + routed_rows: int | None = None, + checkpoint_grad: bool = False, + slot_ref: Any = None, ) -> int: if self._facts is None: return _memory._moe_workspace_bytes( - self, rows, checkpoint_grad=checkpoint_grad, slot_ref=slot_ref + self, + rows, + routed_rows=routed_rows, + checkpoint_grad=checkpoint_grad, + slot_ref=slot_ref, ) group = self._facts["groups"][0 if slot_ref is None else slot_ref] - coefficient, stages = group["gradient" if checkpoint_grad else "forward"] + coefficient, stages, shared = group[ + "gradient" if checkpoint_grad else "forward" + ] return _memory._moe_workspace_from_terms( - rows, (coefficient, tuple(map(tuple, stages))) + rows, (coefficient, tuple(map(tuple, stages)), shared), routed_rows ) def _checkpoint_memory_floor( - self, group_rows: Any, slot_refs: Any = None, gdn_segments: int = 0 + self, + group_rows: Any, + slot_refs: Any = None, + gdn_segments: int = 0, + routed_rows: Any = None, ) -> tuple[int, int]: if self._facts is None: return _memory._checkpoint_memory_floor( - self, group_rows, slot_refs, gdn_segments + self, group_rows, slot_refs, gdn_segments, routed_rows ) return _memory._checkpoint_floor_from_facts( - self, group_rows, slot_refs, gdn_segments, self._facts["checkpoint_layers"] + self, + group_rows, + slot_refs, + gdn_segments, + self._facts["checkpoint_layers"], + routed_rows, ) + def _moe_recompute_covered_for(self, slot_ref: Any) -> bool: + if self._facts is None: + return _memory._moe_recompute_covered_for(self, slot_ref) + if slot_ref is None: + raise ValueError("replayed MoE coverage is frozen per group") + return self._facts["groups"][slot_ref]["moe_covered"] + + def _frozen(self, name: str) -> int: + assert self._facts is not None + return self._facts[name] + + def _moe_checkpoint_state_bytes_per_token(self) -> int: + return self._frozen("moe_checkpoint_state_bytes_per_token") + + def _te_workspace_growth_bytes(self) -> int: + return self._frozen("te_workspace_growth_bytes") + + def _backward_row_state_bytes(self) -> int: + return self._frozen("backward_row_state_bytes") + + def _triton_min_rows(self) -> int: + return self._frozen("triton_min_rows") + def runtime_arguments( self, facts: Any, arguments: dict[str, Any] ) -> dict[str, Any]: @@ -529,6 +731,22 @@ def runtime_arguments( raise ValueError("runtime facts disagree with selected group rows") if arguments.get("hybridep_growth_bytes", 0): raise ValueError("hybridep_runtime_facts_unsupported") + # Capture refuses expert parallelism, where each group's experts see + # its local rows (_plan_group_routed_rows). + routed = tuple(g["rows"] for g in groups) + recorded = arguments.get("group_routed_rows") + if not isinstance(recorded, (list, tuple)) or tuple(recorded) != routed: + raise ValueError("runtime facts disagree with selected routed rows") + traced = arguments.get("head_backward_traced") + if type(traced) is not bool: + raise ValueError("head backward staging is not recorded") + # Live staging needs CP2 and a target backward through a priced head. + if traced and not ( + self._topology_key()[2] == 2 + and facts["head_target_backward"] + and facts["head_vocabulary"] + ): + raise ValueError("head backward staging disagrees with recorded facts") head = max( max( _memory._dense_head_bytes(facts["head_vocabulary"], g["head_rows"]), @@ -560,6 +778,7 @@ def runtime_arguments( return { **arguments, "group_rows": rows, + "group_routed_rows": routed, "slot_refs": tuple(range(len(groups))), "head_workspace_bytes": head, "checkpoint_floor": (retained, workspace), diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 575fafa54..d3540efe4 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -171,6 +171,366 @@ def test_selected_slot_terms_are_replayed_and_frozen(layer, tmp_path): ) +def adapter_report(layer, tmp_path): + """A grouped report whose gradient slot has pending adapter gradients.""" + from test_trainer_rank_adapter_gradient_memory import lora, parameter + from test_trainer_rank_converted_memory import weights + from test_trainer_rank_pending_memory import rank_with_moe + from test_trainer_rank_slot_memory import load_slot + from test_trainer_rank_slot_memory import request as slot_request + + rank, _ = rank_with_moe(weights(layer, 8)) + load_slot(rank, "small", 1) + load_slot(rank, "large", 64) + # Unallocated slot gradients the recompute backward will allocate: small, + # distinct per-layer sizes (the fixture has 40 layers; keep this bounded). + layers = tr._language_model(rank.runtime.model[0]).decoder.layers + assert len(layers) <= 64 + sizes = [16 * (index + 1) for index in range(len(layers))] + assert sum(sizes) * 2 <= 66_560 # BF16 bytes, checked before allocating. + params = [] + for size, block in zip(sizes, layers, strict=True): + params.append(parameter(size)) + block.add_module("adapter", lora(large=[params[-1]])) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward( + [slot_request("small", rows=2), slot_request("large", rows=65, grad=True)], + ensure_slots=False, + ) + original, costs = emitted(rank, plan, tmp_path) + return original, costs, params + + +def test_selected_adapter_gradients_are_replayed_and_frozen(layer, tmp_path): + original, costs, params = adapter_report(layer, tmp_path) + assert costs[0].checkpoint_adapter_gradient > 0 + groups = original["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + assert groups[0]["adapter"] is None + assert groups[1]["adapter"]["name"] == "large" and any( + groups[1]["adapter"]["pending"] + ) + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + assert all(item["matches"] for item in actual["estimates"]) + # Gradients allocated after selection cannot change the replayed answer. + for param in params: + param.grad = torch.zeros_like(param) + assert reports.replay(original) == actual + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["groups"][1][ + "adapter" + ]["pending"][0] += 10**12 + result = reports.replay(changed) + assert not result["estimates"][0]["matches"] + assert ( + result["estimates"][0]["required_bytes"] + > actual["estimates"][0]["required_bytes"] + ) + + +@pytest.mark.parametrize( + "change", + ["gradient", "length", "value", "kind_length", "kindless_pending", "mixed_kinds"], +) +def test_adapter_fact_validation_rejects_forged_input(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + group = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ][1] + adapter = group["adapter"] + if change == "gradient": + group["grad"] = False + group["moe_covered"] = False + elif change == "length": + adapter["pending"] = [0] * 1026 + elif change == "value": + adapter["pending"][0] = 1.5 + elif change == "kind_length": + adapter["kind"] = "k" * 65 + elif change == "kindless_pending": + adapter["kind"] = None + else: + groups = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + groups[0]["grad"] = True + groups[0]["adapter"] = {"kind": None, "name": "base", "pending": []} + # Name the refusal: a forged fact must fail validation, not a later check. + message = ( + "invalid runtime dimension" + if change == "value" + else "invalid adapter gradient facts" + ) + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +def test_recomputed_layer_and_head_stage_facts_are_replayed_and_frozen( + monkeypatch, tmp_path +): + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + original, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + item = original["replay"]["memory_replay"]["estimates"][0] + assert item["runtime_facts"]["groups"][0]["moe_covered"] is True + assert item["arguments"]["head_backward_traced"] is False + actual = reports.replay(original) + assert actual["aggregate"]["matches"] + # The replaying process's settings and TE state do not enter the answer. + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "1") + monkeypatch.setattr(tr, "_TE_CUBLAS_WORKSPACE_BYTES", 0) + assert reports.replay(original) == actual + for field in ( + "moe_checkpoint_state_bytes_per_token", + "te_workspace_growth_bytes", + "backward_row_state_bytes", + "triton_min_rows", + ): + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][field] += ( + 10**9 + ) + result = reports.replay(changed) + assert not result["estimates"][0]["matches"], field + assert ( + result["estimates"][0]["required_bytes"] + > actual["estimates"][0]["required_bytes"] + ), field + # Coverage selects the input-gradient charge and whether the head stages. + changed = deepcopy(original) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["groups"][0][ + "moe_covered" + ] = False + result = reports.replay(changed) + assert not result["estimates"][0]["matches"] + assert ( + result["estimates"][0]["required_bytes"] + != actual["estimates"][0]["required_bytes"] + ) + + +@pytest.mark.parametrize( + "change", + ["shared", "covered_no_grad", "triton_min_rows", "routed_rows", "head_staging"], +) +def test_recomputed_layer_fact_validation_rejects_forged_input(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + item = report["replay"]["memory_replay"]["estimates"][0] + facts = item["runtime_facts"] + if change == "shared": + terms = facts["groups"][1]["gradient"] + terms[2] = terms[0] + 1 + message = "invalid MoE terms" + elif change == "covered_no_grad": + facts["groups"][0]["moe_covered"] = True + message = "invalid MoE recompute coverage" + elif change == "triton_min_rows": + facts["triton_min_rows"] = 64.0 + message = "invalid runtime dimension" + elif change == "routed_rows": + item["arguments"]["group_routed_rows"][1] -= 1 + message = "routed rows" + else: + del item["arguments"]["head_backward_traced"] + message = "head backward staging" + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +@pytest.mark.parametrize("value", [None, 1.5, -1]) +def test_replay_refuses_invalid_recorded_geometry(value, tmp_path): + 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 + ) + # The recomputed mixer multiplies these; refuse before any pricing. + report["replay"]["memory_replay"]["rank"]["geometry"]["kv_channels"] = value + with pytest.raises(ValueError, match="geometry"): + reports.replay(report) + + +@pytest.mark.parametrize("setting", ["0", "-3"]) +def test_non_positive_triton_threshold_is_captured_and_replayed( + setting, monkeypatch, tmp_path +): + # Live pricing accepts it (the fallback chunk then prices no rows). + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", setting) + rank = head_rank() + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, costs = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["triton_min_rows"] == int(setting) + result = reports.replay(report) + assert result["aggregate"]["matches"] + assert result["estimates"][0]["required_bytes"] == costs[0].required + + +def test_recompute_readers_are_unread_without_a_checkpointed_decoder( + monkeypatch, tmp_path +): + from art.trainer_rank import _planner_replay + + # Live pricing reads them only for a checkpointed decoder's recompute; + # capture must not refuse a plan over a setting live never reads. + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "not-a-number") + rank = _rank(monkeypatch) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted(rank, rank._plan_flat_forward([_request(1)]), tmp_path) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["checkpoint_layers"] == 0 + assert all(facts[name] is None for name in _planner_replay._RECOMPUTE_READERS) + assert reports.replay(report)["aggregate"]["matches"] + forged = deepcopy(report) + forged["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "te_workspace_growth_bytes" + ] = 0 + with pytest.raises(ValueError, match="invalid runtime dimension"): + reports.replay(forged) + + +def test_staged_head_needs_recorded_cp2_target_backward(tmp_path): + 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 + ) + # This CP1 plan cannot stage its head; a forged flag must not price it so. + report["replay"]["memory_replay"]["estimates"][0]["arguments"][ + "head_backward_traced" + ] = True + with pytest.raises(ValueError, match="staging disagrees"): + reports.replay(report) + + +def test_staged_cp2_head_is_replayed(monkeypatch, tmp_path): + monkeypatch.setattr( + tr, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + from types import SimpleNamespace + + from art.megatron.context_parallel.types import ParallelTopology + + # CP2 through the stock topology reader (as capture requires), over a CPU + # CP2 topology and planning config. + monkeypatch.setattr(tr.TrainerRank, "_topology_key", lambda self: (1, 1, 2, 1)) + rank = head_rank() + rank._topology = lambda: ParallelTopology(tp=1, cp=2) + for name, value in { + "linear_num_key_heads": 16, + "linear_num_value_heads": 32, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + "params_dtype": torch.bfloat16, + }.items(): + setattr(rank.runtime.provider, name, value) + rank.runtime.model_support_handler = SimpleNamespace( + build_gdn_execution_spec=True, + context_parallel_workload_profile=lambda provider: None, + ) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward([request(512, grad=True)]) + assert plan.signature.topology[2] == 2 + report, costs = emitted(rank, plan, tmp_path) + item = report["replay"]["memory_replay"]["estimates"][0] + assert item["arguments"]["head_backward_traced"] is True + result = reports.replay(report) + assert result["aggregate"]["matches"] + assert result["estimates"][0]["required_bytes"] == costs[0].required + + +@pytest.mark.parametrize("change", ["omit_default", "omit_required", "extra"]) +def test_recorded_geometry_fields_follow_its_schema(change, tmp_path): + 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 + ) + geometry = report["replay"]["memory_replay"]["rank"]["geometry"] + if change == "omit_default": + # Older reports may omit fields that default to zero. + (name,) = [n for n, v in geometry.items() if n == "gdn_conv_kernel" and v == 0] + del geometry[name] + assert reports.replay(report)["aggregate"]["matches"] + return + if change == "omit_required": + del geometry["kv_channels"] + else: + geometry["unknown_width"] = 1 + with pytest.raises(ValueError, match="geometry"): + reports.replay(report) + + +@pytest.mark.parametrize( + "field,value", [("hidden_size", "8"), ("gdn_layers", -1), ("sequence_parallel", 0)] +) +def test_replay_refuses_invalid_recorded_rank_dimensions(field, value, tmp_path): + 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 + ) + report["replay"]["memory_replay"]["rank"][field] = value + with pytest.raises(ValueError, match="rank dimensions"): + reports.replay(report) + + +@pytest.mark.parametrize( + "name", + [ + "_te_workspace_growth_bytes", + "_moe_recompute_covered_for", + "_triton_min_rows", + "_plan_group_routed_rows", + "_plan_head_backward_traced", + "_head_backward_traced", + ], +) +def test_custom_recomputed_layer_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) + + +@pytest.mark.parametrize("layers", [2**10 + 1, 0, 40.0]) +def test_replay_bounds_the_recorded_layer_count(layers, pending_rank, tmp_path): + rank = pending_rank + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(65, grad=True)]), tmp_path + ) + # Replay sizes per-layer tuples from this field; refuse it before costing. + report["replay"]["memory_replay"]["rank"]["num_layers"] = layers + with pytest.raises(ValueError, match="layer count"): + reports.replay(report) + + +def test_custom_adapter_gradient_reader_is_explicitly_incomplete(monkeypatch): + from art.trainer_rank import _planner_replay + + rank = _rank(monkeypatch) + plan = rank._plan_flat_forward([_request(1)]) + original = rank._pending_adapter_gradient_bytes + monkeypatch.setattr( + rank, "_pending_adapter_gradient_bytes", lambda refs: original(refs) + ) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) + + @pytest.mark.parametrize( "change", ["version", "group", "layout", "gdn_segment", "budget"] ) @@ -184,7 +544,7 @@ def test_fact_validation_rejects_inconsistent_or_unbounded_input( ) facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] if change == "version": - facts["version"] = 2 + facts["version"] += 1 elif change == "group": facts["groups"][0]["rows"] += 1 elif change == "layout": diff --git a/tests/unit/test_planner_replay_owner_budget.py b/tests/unit/test_planner_replay_owner_budget.py index 5292cf743..248683dd2 100644 --- a/tests/unit/test_planner_replay_owner_budget.py +++ b/tests/unit/test_planner_replay_owner_budget.py @@ -47,7 +47,7 @@ def test_shared_inventory_preflight_precedes_layout_construction( elif case == "invalid_facts": report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ "version" - ] = 2 + ] += 1 reason = "unsupported runtime facts version" else: monkeypatch.setattr(runtime, "_MAX_INPUT_VALUES", 10) diff --git a/tests/unit/test_planner_runtime_fact_guards.py b/tests/unit/test_planner_runtime_fact_guards.py index aa2002960..9b2a6265f 100644 --- a/tests/unit/test_planner_runtime_fact_guards.py +++ b/tests/unit/test_planner_runtime_fact_guards.py @@ -56,7 +56,7 @@ def test_cumulative_fact_budget_precedes_json_encoding(inventory, monkeypatch): group = facts["groups"][0] group["gdn"] = None if inventory == "stages": - group["forward"] = [0, [[0, 0]] * 4096] + group["forward"] = [0, [[0, 0]] * 4096, 0] group["gradient"] = deepcopy(group["forward"]) facts["groups"] = [deepcopy(group) for _ in range(8)] elif inventory == "slots": diff --git a/tests/unit/test_trainer_rank_adapter_gradient_memory.py b/tests/unit/test_trainer_rank_adapter_gradient_memory.py new file mode 100644 index 000000000..0da30ce0b --- /dev/null +++ b/tests/unit/test_trainer_rank_adapter_gradient_memory.py @@ -0,0 +1,455 @@ +"""Adapter gradients at the recompute backward's peak: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces: a short first wave peaks at layer 0 with +830-900 MB of expert LoRA gradients live, a long one at the last layer. +""" + +from collections.abc import Sequence +import itertools +import random + +import pytest +from test_trainer_rank_checkpoint_memory import rank, requests +import torch + +from art.megatron.lora import LoRA, LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile, _SubforwardCost + +POLICY = LoRASlotRef("checkpoint", "policy") +OTHER = LoRASlotRef("checkpoint", "other") + + +def oracle(pending, boundaries): + """Live bytes beyond the floor while backward recomputes each layer.""" + layers = len(boundaries) + return max( + 0, + *( + sum(pending[index:layers]) + pending[layers] - sum(boundaries[index + 1 :]) + for index in range(layers) + ), + ) + + +def with_pending(monkeypatch, r, pending): + # Like the real method: no slots, nothing pending. + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: tuple(pending) if tuple(refs) else (), + ) + + +@pytest.mark.parametrize( + "pending, boundaries", + [ + # Uniform gradients above the boundaries: layer 0 is the peak. + ([23] * 4 + [0], [4] * 4), + # Uniform gradients below the boundaries: the last layer is the peak. + ([3] * 4 + [0], [10] * 4), + # Attention and GDN layers differ; the peak is interior to neither end. + ([0, 100, 0, 100, 5], [60, 60, 60, 60]), + # Uneven boundaries, gradients outside the decoder live throughout. + ([7, 0, 50, 1, 9], [1, 80, 2, 3]), + ], +) +def test_extra_is_the_worst_backward_layer_over_real_sizes( + monkeypatch, pending, boundaries +): + r = rank() + with_pending(monkeypatch, r, pending) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) == oracle( + pending, boundaries + ) + + +def test_extra_matches_every_layer_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(0) + for _ in range(200): + layers = generator.randint(1, 12) + pending = [ + generator.choice((0, generator.randint(0, 50))) for _ in range(layers + 1) + ] + boundaries = [generator.randint(0, 40) for _ in range(layers)] + with_pending(monkeypatch, r, pending) + extra = r._checkpoint_adapter_gradient_bytes(((POLICY, boundaries),)) + assert extra == oracle(pending, boundaries) + + +def test_no_pending_gradients_add_nothing(monkeypatch): + r = rank() + with_pending(monkeypatch, r, ()) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [4] * 40),)) == 0 + + +def test_a_layer_count_mismatch_prices_no_extra(monkeypatch): + r = rank() + with_slot_pending( + monkeypatch, r, {(POLICY,): [5] * 4 + [0], (OTHER,): [9] * 3 + [0]} + ) + # Pending gradients and boundaries must describe the same decoder layers; + # otherwise the whole term is left out rather than misplaced. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4),)) > 0 + assert r._checkpoint_adapter_gradient_bytes(((OTHER, [1] * 4),)) == 0 + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, [1] * 4), (OTHER, [1] * 4))) == 0 + ) + + +def sequential_oracle(chains): + """Worst live bytes beyond the floor, one group's backward after another. + + While a group recomputes layer i, groups run before it hold all their + gradients and none of their boundaries; groups yet to run hold all their + boundaries; the running group releases its boundaries above i. + """ + worst = 0 + for order in itertools.permutations(chains): + for position, (pending, boundaries) in enumerate(order): + done = order[:position] + layers = len(boundaries) + for index in range(layers): + gradients = ( + sum(sum(p) for p, _ in done) + + pending[layers] + + sum(pending[index:layers]) + ) + released = sum(sum(b) for _, b in done) + sum(boundaries[index + 1 :]) + worst = max(worst, gradients - released) + return worst + + +def with_slot_pending(monkeypatch, r, by_slot): + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: by_slot.get(tuple(refs), ()), + ) + + +def test_gradient_groups_run_their_backward_one_after_another(monkeypatch): + r = rank() + # A short policy group beside a long group of another slot: whichever + # runs first, the other keeps its boundaries (or its gradients) meanwhile. + policy = ([30] * 4 + [0], [2] * 4) + other = ([5] * 4 + [1], [40] * 4) + with_slot_pending(monkeypatch, r, {(POLICY,): policy[0], (OTHER,): other[0]}) + extra = r._checkpoint_adapter_gradient_bytes( + ((POLICY, policy[1]), (OTHER, other[1])) + ) + assert extra == sequential_oracle([policy, other]) + # One chain with every group's boundaries released together would have + # priced far less. + combined = [a + b for a, b in zip(policy[0], other[0])] + assert extra > oracle(combined, [a + b for a, b in zip(policy[1], other[1])]) + # A base-model group owns no gradients, but its boundaries stay live + # while the policy group runs first. + base = ([0] * 5, [40] * 4) + assert ( + r._checkpoint_adapter_gradient_bytes(((POLICY, policy[1]), (None, base[1]))) + == sequential_oracle([policy, base]) + == oracle(*policy) + ) + + +def test_sequential_groups_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(1) + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(5)] + for _ in range(200): + layers = generator.randint(1, 6) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 5)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + if generator.random() < 0.25: + slot, pending = None, [0] * (layers + 1) + else: + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + + +def test_many_gradient_groups_price_exactly(monkeypatch): + r = rank() + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(6)] + # Only groups whose gradients outweigh their boundaries raise another + # group's peak by having run first. + chains = [([3, 1, 2], [100, 100])] * 3 + [([90, 40, 5], [2, 1])] * 3 + with_slot_pending( + monkeypatch, r, {(slot,): pending for slot, (pending, _) in zip(slots, chains)} + ) + groups = [(slot, boundaries) for slot, (_, boundaries) in zip(slots, chains)] + assert r._checkpoint_adapter_gradient_bytes(groups) == sequential_oracle(chains) + assert sequential_oracle(chains) < sum(sum(pending) for pending, _ in chains) + + +def lora(**slots: list[torch.nn.Parameter]) -> LoRA: + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + module._slot_keys = {} + by_name = { + LoRASlotRef("checkpoint", name): params for name, params in slots.items() + } + module.lora_slot_params = lambda ref: by_name.get(ref, []) # type: ignore[method-assign] + return module + + +def parameter(elements: int, *, dtype=torch.bfloat16) -> torch.nn.Parameter: + return torch.nn.Parameter(torch.zeros(elements, dtype=dtype)) + + +def adapter_rank( + layers: Sequence[torch.nn.Module], outside: torch.nn.Module +) -> TrainerRank: + r = rank() + block = r.runtime.model[0].decoder + block.layers = torch.nn.ModuleList(layers) + r.runtime.model[0].head = outside + return r + + +def test_pending_gradients_count_local_unallocated_slot_parameters(): + shared = parameter(10) + allocated = parameter(1000) + allocated.grad = torch.zeros_like(allocated) + master = parameter(1000) + setattr(master, "main_grad", torch.zeros(1000)) + frozen = parameter(1000) + frozen.requires_grad_(False) + layers = [ + lora(policy=[parameter(3), shared]), + lora(policy=[parameter(5)], other=[parameter(1000)]), + lora(policy=[allocated, master, frozen]), + lora(policy=[shared, parameter(4, dtype=torch.float32)]), + ] + r = adapter_rank(layers, lora(policy=[parameter(6)])) + pending = r._pending_adapter_gradient_bytes([POLICY]) + # BF16 bytes per layer; the shared parameter counts once, at its highest + # layer; allocated, main-grad and frozen parameters are not pending; the + # parameter outside the decoder is live throughout. + assert pending == (3 * 2, 5 * 2, 0, (10 + 0) * 2 + 4 * 4, 6 * 2) + assert r._pending_adapter_gradient_bytes([POLICY, OTHER]) == ( + 3 * 2, + 5 * 2 + 1000 * 2, + 0, + 10 * 2 + 4 * 4, + 6 * 2, + ) + assert r._pending_adapter_gradient_bytes([LoRASlotRef("checkpoint", "x")]) == () + assert r._pending_adapter_gradient_bytes([]) == () + + +def test_a_parameter_the_head_also_uses_is_live_throughout(): + tied = parameter(7) + layers = [lora(policy=[tied, parameter(3)]), lora(policy=[parameter(5)])] + r = adapter_rank(layers, lora(policy=[tied])) + # The head's backward runs before the decoder's and allocates it first. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 5 * 2, 7 * 2) + + +def test_a_module_the_head_also_uses_is_live_throughout(): + shared = lora(policy=[parameter(7)]) + r = adapter_rank([shared, torch.nn.Module()], shared) + assert r._pending_adapter_gradient_bytes([POLICY]) == (0, 0, 7 * 2) + + +def test_base_model_groups_own_no_adapter_gradients(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) + n, out, signature, groups, head = values + values = dict( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=((33, True), (34, True), (4096, False)), + slot_refs=(POLICY, LoRASlotRef("checkpoint", None), None), + head_workspace_bytes=head, + ) + both = r._subforward_cost(**values) + # The base group's slot has no adapter, so a split with or without it + # still shares one slot's gradients. + assert both.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + # But its boundaries stay live while the policy group's backward runs. + boundary = 2048 * 2 * 40 + expected = sequential_oracle( + [(pending, [33 * 2048 * 2] * 40), ([0] * 41, [34 * 2048 * 2] * 40)] + ) + assert ( + both.checkpoint_adapter_gradient + == expected + == oracle(pending, [33 * 2048 * 2] * 40) + ) + assert expected > oracle(pending, [(33 + 34) * boundary // 40] * 40) + assert r._estimate_required_memory_bytes_from_values(**values) == both.required + + +def test_other_checkpoint_parameters_are_live_throughout(): + from art.trainer_rank._impl import _CheckpointSlot + + adapter = parameter(3) + custom = parameter(11) + r = adapter_rank([lora(policy=[adapter])], torch.nn.Module()) + r._checkpoint_slots["policy"] = _CheckpointSlot(params=(adapter, custom)) + # A custom object's parameter has no decoder position; LoRA ones keep theirs. + assert r._pending_adapter_gradient_bytes([POLICY]) == (3 * 2, 11 * 2) + + +def test_a_step_with_allocated_gradients_prices_no_extra(): + params = [parameter(100) for _ in range(4)] + r = adapter_rank([lora(policy=[p]) for p in params], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == (200, 200, 200, 200, 0) + for p in params: + p.grad = torch.zeros_like(p) + # Later waves of the step find them in the availability baseline. + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def test_slotless_lora_modules_hold_no_slot_parameters(): + module = LoRA.__new__(LoRA) + torch.nn.Module.__init__(module) + r = adapter_rank([module], torch.nn.Module()) + assert r._pending_adapter_gradient_bytes([POLICY]) == () + + +def priced(r, values, slot_refs): + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + ) + + +def test_cost_and_estimate_charge_the_extra_while_gradients_are_pending(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + retained, workspace = r._checkpoint_memory_floor(groups) + boundary = retained // 40 + pending = [23 * 2**20] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + assert extra > 0 + cost = priced(r, values, (POLICY, None)) + assert cost.checkpoint_adapter_gradient == extra + assert cost.checkpoint_adapter_gradient_slots == '[["checkpoint", "policy"]]' + assert cost.required == int( + (out + retained + workspace + COLD + retained + extra) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Only gradient groups' slots count; no slot, no extra. + assert priced(r, values, (None, POLICY)).checkpoint_adapter_gradient == 0 + with_pending(monkeypatch, r, ()) + plain = priced(r, values, (POLICY, None)) + assert plain.checkpoint_adapter_gradient == 0 + assert plain.required == int((out + 2 * retained + workspace + COLD) * 1.1) + # A learned profile drops the first-execution transients, not the extra. + with_pending(monkeypatch, r, pending) + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + assert priced(r, values, (POLICY, None)).required == int( + (out + 2 * retained + workspace + extra) * 1.1 + ) + + +def child(extra: int, slots: str, workspace: int = 10) -> _SubforwardCost: + return _SubforwardCost( + required=int((1 + workspace + 1 + extra) * 1.1), + retained=0, + checkpoint_retained=1, + checkpoint_workspace=workspace, + checkpoint_input_gradient=1, + checkpoint_adapter_gradient=extra, + checkpoint_adapter_gradient_slots=slots, + ) + + +def test_split_charges_shared_gradients_once_and_distinct_slots_each(): + shared = [child(500, "a"), child(300, "a")] + assert TrainerRank._split_required_memory(shared) == int((2 + 2 + 10 + 500) * 1.1) + distinct = [child(500, "a"), child(300, "b")] + assert TrainerRank._split_required_memory(distinct) == int((2 + 2 + 10 + 800) * 1.1) + # A child with no pending gradients does not change the shared charge. + assert TrainerRank._split_required_memory([child(500, "a"), child(0, "")]) == int( + (2 + 2 + 10 + 500) * 1.1 + ) + + +def test_cheap_estimate_defers_while_a_gradient_slot_has_pending_gradients( + monkeypatch, +): + r = rank() + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr( + r, + "_resolve_slot_ref", + lambda request, checkpoint: POLICY if not request.no_grad else None, + ) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is None + assert seen == [(POLICY,)] + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", lambda refs: ()) + assert r._estimate_flat_forward(requests(67, 4096)) is not None + + +def test_a_base_only_gradient_group_prices_and_defers_nothing(monkeypatch): + r = rank() + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, head = values + base = LoRASlotRef("checkpoint", None) + with_pending(monkeypatch, r, [23 * 2**20] * 40 + [0]) + cost = priced(r, values, (base, None)) + # The base model has no adapter: no extra, no slot identity. + assert cost.checkpoint_adapter_gradient == 0 + assert cost.checkpoint_adapter_gradient_slots == "" + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(base, None), + head_workspace_bytes=head, + ) + assert estimate == cost.required + # Nor does the cheap estimate defer to the exact plan for it. + monkeypatch.setattr(r, "_ensure_checkpoint_slots_for", lambda *a, **k: None) + monkeypatch.setattr(r, "_resolve_slot_ref", lambda request, checkpoint: base) + seen: list[tuple[LoRASlotRef, ...]] = [] + + def pending(refs): + seen.append(tuple(refs)) + return (1,) * 41 + + monkeypatch.setattr(r, "_pending_adapter_gradient_bytes", pending) + assert r._estimate_flat_forward(requests(67, 4096)) is not None + assert seen == [] diff --git a/tests/unit/test_trainer_rank_admission_inputs.py b/tests/unit/test_trainer_rank_admission_inputs.py index a79d198e1..3481bfd6d 100644 --- a/tests/unit/test_trainer_rank_admission_inputs.py +++ b/tests/unit/test_trainer_rank_admission_inputs.py @@ -32,6 +32,7 @@ def assert_plan_values(rank, plan, values): assert values["signature"] == plan.signature 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["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_gradient_memory.py b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py index d6d9ded56..c27ca5e01 100644 --- a/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_gradient_memory.py @@ -1,14 +1,26 @@ -"""Partial input-gradient extents: CPU admission math, not peak/overlap proof.""" +"""Input-gradient extents: CPU admission math, not peak/overlap proof. + +Where the recomputed MoE stage is priced, the floor charges the one incoming +gradient live at the last layer's recompute peak; elsewhere it keeps one +gradient per saved boundary. +""" from dataclasses import replace import pytest from test_trainer_rank_checkpoint_memory import price, rank, requests -from test_trainer_rank_moe_memory import layer # noqa: F401 -from test_trainer_rank_pending_memory import full_requests, pending_rank # noqa: F401 +from test_trainer_rank_moe_memory import _enclosing_moe, layer # noqa: F401 +from test_trainer_rank_pending_memory import ( # noqa: F401 + full_requests, + pending_rank, + rank_with_moe, +) import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import Unset, _MemoryProfile, _SplitForwardPlan @@ -25,25 +37,93 @@ def test_pending_cold_peak_does_not_become_forward_retention(pending_rank): r = pending_rank plan = r._plan_flat_forward(full_requests()) cost = r._plan_cost(plan) - gradient = 8 * 6330 * 40 * 2048 * 2 + boundaries = 8 * 6330 * 40 * 2048 * 2 + gradient = 8 * 6330 * 2048 * 2 assert cost.checkpoint_input_gradient == gradient - # Exact previous cold estimate, including outputs and its one safety factor. + # Exact cold estimate, including outputs and its one safety factor. assert cost.retained == 23102959299 - assert cost.required == int((plan.output_bytes + 2 * gradient + 12705630112) * 1.1) + assert cost.required == int( + (plan.output_bytes + boundaries + gradient + cost.checkpoint_workspace) * 1.1 + ) assert r._memory_check(plan).estimated_required_bytes == cost.required profile(r, plan) warm = r._plan_cost(plan) - assert warm.retained == int((plan.output_bytes + gradient) * 1.1) - assert warm.required == cost.required + assert warm.retained == int((plan.output_bytes + boundaries) * 1.1) + # Profiled: no first-execution transients. + assert warm.required == int( + (plan.output_bytes + boundaries + gradient + cost.checkpoint_workspace - COLD) + * 1.1 + ) + + +BOUNDARY_GRADIENTS = 8 * 6330 * 40 * 2048 * 2 + + +def test_single_gradient_needs_every_layer_priced(layer): + # One MoE layer among 40 dense stand-ins: the floor does not price the dense + # layers' recompute, so one gradient per boundary stays. + r = rank_with_moe(_enclosing_moe(layer), stand_in=False)[0] + assert r._moe_gradient_enclosed == (True,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_single_gradient_needs_the_fc1_stage(layer): + # Without permute fusion the FC1 stage is not enclosed: a positive FC2-only + # coefficient does not cover recompute. + moe = _enclosing_moe(layer) + moe.config.moe_permute_fusion = False + r = rank_with_moe(moe)[0] + assert r._checkpoint_moe_bytes_per_token() > 0 + assert r._moe_gradient_enclosed == (False,) and not r._moe_recompute_covered + plan = r._plan_flat_forward(full_requests()) + assert r._plan_cost(plan).checkpoint_input_gradient == BOUNDARY_GRADIENTS + + +def test_a_slot_that_loses_moe_coverage_keeps_boundary_gradients( + pending_rank, monkeypatch +): + from art.megatron.lora import LoRASlotRef + from art.trainer_rank import _impl + + r = pending_rank + ref = LoRASlotRef(kind="checkpoint", name="policy") + groups = ((100, True),) + assert r._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", lambda *a, **k: 0) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + + def covered(*args, enclosed, **kwargs): + enclosed.extend([True] * len(r._moe_gradient_enclosed)) + return 1 + monkeypatch.setattr(_impl, "_moe_output_bytes_per_token", covered) + assert r._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + +@pytest.mark.parametrize("cp", [1, 2, 4]) +def test_one_moe_gradient_only_up_to_the_traced_cp2(pending_rank, monkeypatch, cp): + """Above CP2 a rank runs more remote attention stages than the mixer's CP2 + allowance prices; the per-boundary gradient allowance must cover them.""" + r = pending_rank + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, cp, 1)) + groups = ((100, True),) + one, every = 100 * 2048 * 2, 100 * 40 * 4096 + assert r._checkpoint_input_gradient_bytes(groups) == (one if cp <= 2 else every) + + +@pytest.mark.parametrize("moe", [True, False]) @pytest.mark.parametrize("rows", [1, 67, 1024]) -def test_attention_only_extent_scales_with_gradient_rows(rows): +def test_attention_only_extent_scales_with_gradient_rows(rows, moe): r = rank() + if not moe: + # Without a priced MoE stage, one gradient per boundary stays: it also + # covers the dense MLP and other recompute work the floor omits. + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 values = r._estimate_flat_forward(requests(rows, 4096)) cost = price(r, values) assert r._gdn_layers == 0 - assert cost.checkpoint_input_gradient == rows * 40 * 2048 * 2 + assert cost.checkpoint_input_gradient == rows * (40 if not moe else 1) * 2048 * 2 assert cost.required >= int( ( cost.checkpoint_retained @@ -59,10 +139,19 @@ def test_gradient_is_not_absorbed_by_larger_head_workspace(): n, out, sig, groups, _head = r._estimate_flat_forward(requests(67, 4096)) head = 10**10 cost = price(r, (n, out, sig, groups, head)) - gradient = 67 * 40 * 2048 * 2 - assert cost.checkpoint_workspace == head - assert cost.required == int((out + head + 2 * gradient) * 1.1) - assert cost.retained == int((out + head + gradient) * 1.1) + boundaries = 67 * 40 * 2048 * 2 + gradient = 67 * 2048 * 2 + # The head stage: two more gradient-row terms, each row's RoPE and index + # state, and TE's first-GEMM workspaces. + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) + assert cost.checkpoint_workspace == stage + COLD + assert cost.required == int((out + stage + COLD + boundaries + gradient) * 1.1) + assert cost.retained == int((out + head + boundaries) * 1.1) r._memory_profiles[sig] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=n, @@ -111,11 +200,7 @@ def test_split_sums_all_gradient_children_outside_workspace_max(): profile(r, child) costs = [r._plan_cost(child) for child in children] split = _SplitForwardPlan(tuple(children), ((0,), (1,), (2,)), 3) - assert [c.checkpoint_input_gradient for c in costs] == [ - 17 * 40 * 4096, - 29 * 40 * 4096, - 0, - ] + assert [c.checkpoint_input_gradient for c in costs] == [17 * 4096, 29 * 4096, 0] expected = int( ( sum(c.checkpoint_retained + c.checkpoint_input_gradient for c in costs) @@ -157,7 +242,7 @@ def test_lower_bound_profile_cliff_preserves_separate_peak_component(): lower = r._split_chunk_lower_cost( req, tuple(q.input_tokens for q in req), checkpoint=Unset ) - assert lower.checkpoint_input_gradient == 128 * 40 * 4096 + assert lower.checkpoint_input_gradient == 128 * 4096 assert lower.required <= r._plan_cost(full).required assert lower.checkpoint_retained == full.output_bytes + 128 * 40 * 4096 @@ -199,16 +284,18 @@ def test_split_priority_subtracts_only_uncovered_gradient_peak(fully_masked): r = rank() plan = r._plan_flat_forward(requests(17, 19)) cold = r._plan_cost(plan) + # Once profiled, the static estimate has no first-execution transients. + profile(r, plan) + static = r._plan_cost(plan).required + assert static < cold.required # Place a real learned peak between the two static estimates, or above both. - measured = ( - cold.required + 10**7 if fully_masked else (cold.retained + cold.required) / 2 - ) + measured = static + 10**7 if fully_masked else (cold.retained + static) / 2 rate = (measured / 1.1 - plan.output_bytes) / plan.packed_tokens profile(r, plan, rate=rate) cost = r._plan_cost(plan) old_required = int((plan.output_bytes + int(plan.packed_tokens * rate)) * 1.1) assert cold.retained < old_required - assert cost.required == max(cold.required, old_required) + assert cost.required == max(static, old_required) assert ( cost.ephemeral - cost.checkpoint_peak_increment == old_required - cost.retained ) @@ -220,3 +307,66 @@ def test_split_priority_subtracts_only_uncovered_gradient_peak(fully_masked): < cost.checkpoint_peak_increment < int(cost.checkpoint_input_gradient * 1.1) ) + + +def test_one_moe_gradient_still_prices_pending_adapter_gradients( + pending_rank, monkeypatch +): + from art.megatron.lora import LoRASlotRef + from art.trainer_rank import _gdn_memory + + r = pending_rank + plan = r._plan_flat_forward(full_requests()) + groups = r._plan_group_rows(plan) + retained, _ = r._checkpoint_memory_floor(groups) + # The MoE stage covers recompute: one incoming gradient, no per-boundary + # slack left to absorb gradients the backward allocates. + assert r._checkpoint_input_gradient_bytes(groups) < retained + pending = (23 * 2**20,) * 40 + (0,) + monkeypatch.setattr( + r, + "_pending_adapter_gradient_bytes", + lambda refs: pending if tuple(refs) else (), + ) + # A named slot keeping the constructor's full MoE coverage. This fixture + # has no slot tables, so price the MoE stage with the constructor's adapters. + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + workspace = r._moe_workspace_bytes + monkeypatch.setattr( + r, + "_moe_workspace_bytes", + lambda rows, **kwargs: workspace(rows, **{**kwargs, "slot_ref": None}), + ) + values = dict( + packed_tokens=plan.packed_tokens, + output_bytes=plan.output_bytes, + signature=plan.signature, + logical_tokens=plan.active_logical_tokens, + gdn_segments=plan.grad_segment_count, + group_rows=groups, + group_routed_rows=r._plan_group_routed_rows(plan), + head_workspace_bytes=r._plan_head_workspace_bytes(plan), + checkpoint_floor=_gdn_memory.plan_floor(r, plan), + retained_tokens=r._plan_retained_tokens(plan), + ) + slotless = (None,) * len(groups) + named = (LoRASlotRef("checkpoint", "policy"),) * len(groups) + base = r._subforward_cost(**values, slot_refs=slotless) + cost = r._subforward_cost(**values, slot_refs=named) + boundary = retained // 40 + extra = max(0, *(sum(pending[i:40]) - boundary * (39 - i) for i in range(40))) + assert base.checkpoint_adapter_gradient == 0 + assert cost.checkpoint_adapter_gradient == extra > 0 + assert cost.required == int( + ( + cost.checkpoint_retained + + cost.checkpoint_workspace + + cost.checkpoint_input_gradient + + extra + ) + * 1.1 + ) + assert ( + r._estimate_required_memory_bytes_from_values(**values, slot_refs=named) + == cost.required + ) diff --git a/tests/unit/test_trainer_rank_checkpoint_memory.py b/tests/unit/test_trainer_rank_checkpoint_memory.py index c6e33ee9d..78fcb7f73 100644 --- a/tests/unit/test_trainer_rank_checkpoint_memory.py +++ b/tests/unit/test_trainer_rank_checkpoint_memory.py @@ -8,7 +8,12 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank -from art.trainer_rank._impl import Unset, _ForwardRefusal, _MemoryProfile +from art.trainer_rank._impl import ( + _TE_CUBLAS_WORKSPACE_BYTES, + Unset, + _ForwardRefusal, + _MemoryProfile, +) def rank(): @@ -53,6 +58,8 @@ def rank(): ) result._moe_output_bytes_per_token = 188416 result._moe_checkpoint_grad_bytes_per_token = 188416 + # Stands in for Qwen3.6-35B-A3B, whose MoE stage prices every layer's recompute. + result._moe_recompute_covered = True return result @@ -105,7 +112,9 @@ def test_required_and_learned_retained_use_max_not_sum(): ) cost = price(r, values) old = max(n * 2048 * 2 * 14, n * 188416) - assert cost.required == int((out + max(old, 2 * retained + work)) * 1.1) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == 1024 * 2048 * 2 # The incoming gradient only. + assert cost.required == int((out + max(old, retained + gradient + work)) * 1.1) assert cost.retained == int((out + retained) * 1.1) r._memory_profiles[sig] = replace( r._memory_profiles[sig], @@ -135,7 +144,7 @@ def test_per_group_padding_precedes_gradient_filter(): assert values[0] == 40 and values[3] == ((16, True), (24, False)) assert r._checkpoint_memory_floor(values[3]) == ( 16 * 40 * 4096, - 24 * (188416 + 4 * 2048 * 2), + 24 * (188416 + 4 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -193,11 +202,18 @@ def test_topology_revalidated(axis): @pytest.mark.parametrize("rows", [(10, True), (11, False)]) @pytest.mark.parametrize("cp", [2, 4]) def test_cp_floor_prices_rank_rows(cp, rows): - # Callers pass rows on the most loaded CP rank; the per-row floor matches CP1. + # Callers pass rows on the most loaded CP rank; the per-row floor matches + # CP1, except that CP attention also keeps its stage buffers while a + # gradient group recomputes the layer. r = rank() single = r._checkpoint_memory_floor((rows,)) r._topology_key = lambda: (1, 1, cp, 1) - assert r._checkpoint_memory_floor((rows,)) == single != (0, 0) + retained, workspace = r._checkpoint_memory_floor((rows,)) + assert retained == single[0] and single != (0, 0) + count, grad = rows + # No attention geometry in this stub: Q and KV widths fall back to hidden. + stage = count * 2 * (3 * 2048 + 2 * 2048) if grad else 0 + assert workspace == single[1] + stage @pytest.mark.parametrize("share", [lambda n: -(-n // 2), lambda n: n * 3 // 4]) @@ -318,7 +334,8 @@ def test_split_keeps_complete_order_and_checks_each_new_subforward(): logical_per_packed=1, retained_compute_bytes_per_token=1, ) - limit = 250_000_000 + # Just below the unsplit plan's requirement: two subforwards must fit. + limit = r._plan_cost(flat).required - 1 used = 0 r._available_memory_bytes = lambda: limit - used result = r._find_admissible_forward(req, checkpoint=Unset, refusal_prefix="test") @@ -424,7 +441,8 @@ def test_no_grad_enclosure_uses_max_group_and_affine_stage(): mixed = ((3, True), (11, False)) assert r._checkpoint_memory_floor(mixed) == ( 3 * 40 * 2048 * 2, - max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2), + max(3 * 188416, max(11 * 188416, 11 + 1_000_000) + 4 * 11 * 2048 * 2) + + _TE_CUBLAS_WORKSPACE_BYTES, ) @@ -497,3 +515,155 @@ def test_no_grad_enclosure_config_guard(field, value): r = rank() setattr(r.runtime.model[0].decoder.config, field, value) assert r._checkpoint_memory_floor(((11, False),)) == (0, 0) + + +def test_routed_rows_move_only_the_routed_moe_part(): + # A CP2/EP2 real-data trace: the busiest CP rank held 52,480 rows, while + # HybridEP dispatched 8 x 96,794 pairs per layer across both ranks, so a + # balanced rank receives 48,397 rows' pairs. Boundaries, the mixer and the + # shared expert stay on the local rows. + r = rank() + r._topology_key = lambda: (1, 1, 2, 1) + r._moe_gradient_shared_bytes = 8192 + local, routed = 52480, 48397 + retained, workspace = r._checkpoint_memory_floor(((local, True),)) + assert r._checkpoint_memory_floor(((local, True),), None, routed_rows=(local,)) == ( + retained, + workspace, + ) + 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) + r._moe_gradient_shared_bytes = 188417 + with pytest.raises(ValueError, match="shared-expert"): + r._moe_workspace_bytes(10, checkpoint_grad=True) + + +def qwen36_attention(r): + # Qwen3.6-35B-A3B attention: 16 heads and 2 query groups of 256, gated. + r._geometry = replace( + r._geometry, num_attention_heads=16, num_query_groups=2, kv_channels=256 + ) + r._attention_output_gate = True + return r + + +@pytest.mark.parametrize( + "cp,expected", + # Traces measured 66 KB per token at CP1 and 94 KB at CP2. CP4 reuses the + # CP2 stage allowance; it is not measured. + [(1, 2 * (2 * 2048 + 7 * 4096 + 3 * 512)), (2, 95232), (4, 95232)], +) +def test_recomputed_attention_is_priced_beside_the_moe_stage(cp, expected): + r = qwen36_attention(rank()) + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + moe = r._moe_workspace_bytes(10, checkpoint_grad=True) + # Beside the mixer: the recomputed layer's residual and pre-MLP norm rows, + # and Transformer Engine's cuBLAS workspaces. + assert ( + r._checkpoint_memory_floor(((10, True),))[1] + == moe + 10 * (expected + 2 * 2048 * 2) + _TE_CUBLAS_WORKSPACE_BYTES + ) + + +@pytest.mark.parametrize( + "cp,gdn_layers,expected", + [ + # Hybrid: GDN (82 KB) is larger at CP1, CP attention (95 KB) at CP2. + (1, 30, 81920), + (2, 30, 95232), + # GDN only: CP exchanges add a hidden-width and a value-width row. + (1, 40, 81920), + (2, 40, 94208), + ], +) +def test_larger_recomputed_mixer_is_priced(cp, gdn_layers, expected): + r = qwen36_attention(rank()) + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=32, + gdn_value_head_dim=128, + ) + r._gdn_layers = gdn_layers + r._topology_key = lambda: (1, 1, cp, 1) + assert r._recomputed_mixer_bytes_per_token() == expected + + +def test_no_grad_groups_do_not_recompute_a_mixer(): + r = qwen36_attention(rank()) + moe = r._moe_workspace_bytes(10) + assert r._checkpoint_memory_floor(((10, False),)) == (0, moe + 4 * 10 * 2048 * 2) + + +def test_ungated_attention_prices_fewer_projections(): + r = qwen36_attention(rank()) + r._attention_output_gate = False + assert r._recomputed_mixer_bytes_per_token() == 2 * (2 * 2048 + 5 * 4096 + 3 * 512) + r._topology_key = lambda: (1, 1, 2, 1) + assert r._recomputed_mixer_bytes_per_token() == 2 * ( + 2 * 2048 + 5 * 4096 + 3 * 512 + 3 * 4096 + 2 * 512 + ) + + +@pytest.mark.parametrize("cp", [1, 2]) +def test_gdn_width_follows_hidden_key_and_value_separately(cp): + # Hidden differs from the key width, as in Qwen3.5-27B: CP exchanges + # carry hidden-width inputs and value-width outputs. + r = rank() + r._hidden_size = r.runtime.model[0].decoder.config.hidden_size = 5120 + r._geometry = replace( + r._geometry, + gdn_key_heads=16, + gdn_key_head_dim=128, + gdn_value_heads=48, + gdn_value_head_dim=128, + ) + r._gdn_layers = r._num_layers + r._topology_key = lambda: (1, 1, cp, 1) + key, value = 16 * 128, 48 * 128 + # q and k are l2-normalized after expansion to the 48 value heads. + width = 5120 + 2 * key + 2 * 48 * 128 + 6 * value + 64 * 48 + if cp > 1: + width += 5120 + value + assert r._recomputed_mixer_bytes_per_token() == 2 * width + moe = r._moe_workspace_bytes(7, checkpoint_grad=True) + assert ( + r._checkpoint_memory_floor(((7, True),))[1] + == moe + 7 * 2 * (width + 2 * 5120) + _TE_CUBLAS_WORKSPACE_BYTES + ) + + +@pytest.mark.parametrize("ep,routing", [(1, 3332), (2, 3600)]) +def test_moe_state_beside_the_recomputed_stage(ep, routing): + # Qwen3.6-35B-A3B: 256 experts, 512-wide shared expert. Traces measured + # 3,332 (EP1) and 3,593 (EP2) bytes of routing state per local token. + r = rank() + r._geometry = replace(r._geometry, moe_experts=256, moe_shared_expert_ffn=512) + r._parallel_shape = replace(r._parallel_shape, ep=ep) + assert r._moe_checkpoint_state_bytes_per_token() == routing + 3 * 512 * 2 + r._geometry = replace(r._geometry, moe_experts=0) + assert r._moe_checkpoint_state_bytes_per_token() == 0 + + +def test_te_workspaces_are_growth_until_allocated(monkeypatch): + gemm = pytest.importorskip("transformer_engine.pytorch.cpp_extensions.gemm") + r = rank() + entries = [0] + + class Cached: + def cache_info(self): + return SimpleNamespace(currsize=entries[0]) + + monkeypatch.setattr(gemm, "get_cublas_workspace", Cached()) + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 1 # Plain GEMM only: the grouped streams are still to come. + assert r._te_workspace_growth_bytes() == _TE_CUBLAS_WORKSPACE_BYTES + entries[0] = 2 + assert r._te_workspace_growth_bytes() == 0 diff --git a/tests/unit/test_trainer_rank_converted_memory.py b/tests/unit/test_trainer_rank_converted_memory.py index a9c809163..fe1b47372 100644 --- a/tests/unit/test_trainer_rank_converted_memory.py +++ b/tests/unit/test_trainer_rank_converted_memory.py @@ -1,16 +1,21 @@ """Source-derived affine routed-expert stages; no complete backward/compiled bound.""" +import math from types import SimpleNamespace from typing import Any, cast import pytest -from test_trainer_rank_moe_memory import _enclosing_moe, _rank +from test_trainer_rank_moe_memory import _enclosing_moe, _hybridep, _rank from test_trainer_rank_moe_memory import layer as layer from test_trainer_rank_pending_memory import module, rank_with_moe import torch from art.trainer_rank import ForwardInput, _gdn_memory -from art.trainer_rank._impl import _expert_lora_weight_storage +from art.trainer_rank._impl import ( + _expert_lora_weight_storage, + _moe_output_bytes_per_token, +) +from art.trainer_rank._planner_cost import ParallelShape def weights(layer: Any, rank: int, *, fc1: bool = True, dtype=torch.bfloat16): @@ -96,9 +101,19 @@ def test_actual_plan_cost_and_admission(layer, rank_value, grad, output): retained, workspace = rank._checkpoint_memory_floor(rank._plan_group_rows(plan)) pending = _gdn_memory.plan_floor(rank, plan) if grad: - assert workspace == expected(8, rank_value, True) + # The recomputed layer's mixer, its residual and pre-MLP norm rows and + # its MoE routing state stay live beside its MoE stage; the first call + # also allocates TE's cuBLAS workspaces. + mixer = rank._recomputed_mixer_bytes_per_token() + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() + assert mixer > 0 + assert workspace == ( + expected(8, rank_value, True) + + 8 * (mixer + beside) + + rank._te_workspace_growth_bytes() + ) + # The GDN pending floor combines with this one by maximum. assert pending[0] == retained == 8 * 40 * 2048 * 2 - assert pending[1] >= workspace else: assert retained == 0 and pending == (0, 0) assert workspace == expected(8, rank_value, False) + 4 * 8 * 2048 * 2 @@ -124,8 +139,15 @@ def test_reference_and_gradient_keep_distinct_stage_modes(layer, order, rank_val assert set(groups) == {(3, True), (9, False)} retained, workspace = rank._checkpoint_memory_floor(groups) assert retained == 3 * 40 * 2048 * 2 - assert workspace == max( - expected(3, rank_value, True), expected(9, rank_value, False) + 4 * 9 * 2048 * 2 + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() + assert ( + workspace + == max( + expected(3, rank_value, True) + + 3 * (rank._recomputed_mixer_bytes_per_token() + beside), + expected(9, rank_value, False) + 4 * 9 * 2048 * 2, + ) + + rank._te_workspace_growth_bytes() ) assert ( rank._memory_check(plan).estimated_required_bytes @@ -262,6 +284,38 @@ def test_fc1_fixed_weights_do_not_scale_with_topk(layer, topk): ) +def test_hybridep_fc1_stages_hold_one_dispatched_input(layer): + # FC1's converted stages hold the routed H-wide inputs too: two under the + # EP1 all-to-all, one under HybridEP, over its EP2 allowance rows. + weights(layer, 8) + single: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [layer], + ParallelShape(tp=1, cp=1), + checkpoint_grad=True, + converted_stages=single, + ) + expert = _hybridep(layer, 2) + expert.token_dispatcher.num_local_experts = 128 + sharded: list[tuple[int, int]] = [] + _moe_output_bytes_per_token( + [expert], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=True, + converted_stages=sharded, + ) + assert [stage[0] for stage in single[:2]] == [ + 16 * (2 * 2048 + 2 * 1024 + 8), + 16 * (2 * 2048 + 3 * 1024 + 8), + ] + # 8 x 1.4 routed rows at EP2, rounded up per stage. + assert [stage[0] for stage in sharded[:2]] == [ + math.ceil(8 * 1.4 * 2 * (2048 + 2 * 1024 + 8)), + math.ceil(8 * 1.4 * 2 * (2048 + 3 * 1024 + 8)), + ] + assert [stage[1] for stage in sharded] == [stage[1] for stage in single] + + def test_wide_fc1_sum_is_a_separate_stage(layer): weights(layer, 8) experts = layer.experts diff --git a/tests/unit/test_trainer_rank_head_memory.py b/tests/unit/test_trainer_rank_head_memory.py index 21b7ad071..569037e40 100644 --- a/tests/unit/test_trainer_rank_head_memory.py +++ b/tests/unit/test_trainer_rank_head_memory.py @@ -8,6 +8,9 @@ import torch from art.trainer_rank import ForwardInput +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, Unset, @@ -143,11 +146,19 @@ def test_outputs_retention_and_empirical_peak_are_counted_once(): r = rank() plan = r._plan_flat_forward([request(grad=True)]) retained = 512 * 40 * 2048 * 2 - gradient = 512 * 40 * 2048 * 2 + gradient = 512 * 2048 * 2 # The incoming gradient beside the MoE stage. head = 3 * 512 * 248320 * 2 cost = r._plan_cost(plan) assert cost.retained == int((plan.output_bytes + retained + head) * 1.1) - assert cost.required == int((plan.output_bytes + retained + gradient + head) * 1.1) + # Unprofiled head stage: the head workspace, the final output, saved + # selected rows and hidden-row gradient beside the incoming one, each + # row's RoPE and index state, TE's first-GEMM workspaces and the first + # execution's transients. + state = 512 * r._backward_row_state_bytes() + te = r._te_workspace_growth_bytes() + assert cost.required == int( + (plan.output_bytes + retained + 3 * gradient + state + head + te + COLD) * 1.1 + ) r._memory_profiles[plan.signature] = _MemoryProfile( bytes_per_token=2_000_000, packed_tokens=512, @@ -306,14 +317,21 @@ def test_tied_standard_head_weight_uses_the_same_capacity(): @pytest.mark.parametrize("rows", [128, 512]) -def test_target_backward_refuses_budget_below_logits_and_both_gradients(rows): +def test_target_backward_refuses_budget_below_logits_and_both_gradients( + monkeypatch, rows +): r = rank() + # Isolate the head term from TE's one-time cuBLAS workspace growth, and + # the three target-backward buffers from the small-chunk fallback cover. + r._te_workspace_growth_bytes = lambda: 0 + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "1") plan = r._plan_flat_forward([request(rows, grad=True)]) retained, _ = r._checkpoint_memory_floor(r._plan_group_rows(plan)) - gradient = rows * 40 * 2048 * 2 + gradient = rows * 2048 * 2 + stage = 3 * gradient + rows * r._backward_row_state_bytes() dense = min(rows, 512) * 248320 * 2 - before = int((plan.output_bytes + retained + gradient + 2 * dense) * 1.1) - expected = int((plan.output_bytes + retained + gradient + 3 * dense) * 1.1) + before = int((plan.output_bytes + retained + stage + 2 * dense + COLD) * 1.1) + expected = int((plan.output_bytes + retained + stage + 3 * dense + COLD) * 1.1) r._available_memory_bytes = lambda: (before + expected) // 2 check = r._memory_check(plan) assert check.estimated_required_bytes == expected diff --git a/tests/unit/test_trainer_rank_head_stage_memory.py b/tests/unit/test_trainer_rank_head_stage_memory.py new file mode 100644 index 000000000..695ebe274 --- /dev/null +++ b/tests/unit/test_trainer_rank_head_stage_memory.py @@ -0,0 +1,456 @@ +"""The head's backward and the decoder's recompute peak apart: CPU admission math. + +Qwen3.6-35B-A3B CP2 allocator traces (EP1 and EP2): the head's buffers are +freed before the recompute backward allocates its workspace and adapter +gradients, and TE's first-GEMM workspaces are live at the head's peak. +""" + +from dataclasses import replace +import itertools +import random +from typing import Any + +import pytest +from test_trainer_rank_adapter_gradient_memory import ( + OTHER, + POLICY, + oracle, + priced, + with_pending, + with_slot_pending, +) +from test_trainer_rank_checkpoint_memory import rank, requests + +from art.megatron.lora import LoRASlotRef +from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +from art.trainer_rank._impl import _MemoryProfile + +MiB = 2**20 + + +def covered_rank(monkeypatch): + r = rank() + # The MoE stage covers the policy slot's recompute, as the constructor's. + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + return r + + +def head_oracle(chains): + """Worst adapter bytes live beyond the floor while one group's head runs. + + Any set of the other groups may have run their backward first: each holds + all its gradients and none of its boundaries. The running group holds only + its gradients outside the decoder. + """ + worst = 0 + for index, (pending, _) in enumerate(chains): + others = chains[:index] + chains[index + 1 :] + for count in range(len(others) + 1): + for done in itertools.combinations(others, count): + live = pending[-1] + sum(sum(p) - sum(b) for p, b in done) + worst = max(worst, live) + return worst + + +def test_head_adapter_gradients_match_every_order_for_random_sizes(monkeypatch): + r = rank() + generator = random.Random(2) + slots = [LoRASlotRef("checkpoint", f"slot{index}") for index in range(4)] + for _ in range(200): + layers = generator.randint(1, 5) + chains, groups, by_slot = [], [], {} + for slot in slots[: generator.randint(1, 4)]: + pending = [generator.randint(0, 30) for _ in range(layers + 1)] + boundaries = [generator.randint(0, 30) for _ in range(layers)] + by_slot[(slot,)] = pending + chains.append((pending, boundaries)) + groups.append((slot, boundaries)) + with_slot_pending(monkeypatch, r, by_slot) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == head_oracle( + chains + ) + + +def test_one_group_head_meets_only_its_gradients_outside_the_decoder(monkeypatch): + r = rank() + with_pending(monkeypatch, r, [30] * 4 + [7]) + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),), head=True) == 7 + # The decoder stage still meets its own gradients. + assert r._checkpoint_adapter_gradient_bytes(((POLICY, [2] * 4),)) == oracle( + [30] * 4 + [7], [2] * 4 + ) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [1] * 4 + [0]} + ) + # The policy group may run first and raise the other's head; a group whose + # boundaries outweigh its gradients never raises the policy's. + groups = ((POLICY, [2] * 4), (OTHER, [40] * 4)) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + groups = ((POLICY, [2] * 4), (OTHER, [2] * 4)) + with_slot_pending( + monkeypatch, r, {(POLICY,): [30] * 4 + [7], (OTHER,): [0] * 4 + [0]} + ) + assert r._checkpoint_adapter_gradient_bytes(groups, head=True) == 127 - 8 + + +def staged(r, values, slot_refs): + """``priced`` for a head whose backward is the traced one.""" + n, out, signature, groups, head = values + return r._subforward_cost( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=slot_refs, + head_workspace_bytes=head, + head_backward_traced=True, + ) + + +@pytest.mark.parametrize("head", [0, 64 * MiB, 700 * MiB, 3000 * MiB]) +def test_cost_and_estimate_price_the_larger_stage(monkeypatch, head): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + values = (n, out, signature, groups, head) + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = 67 * 2048 * 2 + boundary = retained // 40 + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [boundary] * 40) + te = r._te_workspace_growth_bytes() + state = 67 * r._backward_row_state_bytes() + decoder = workspace + extra + stage = head + 2 * gradient + state + te if head else 0 + cost = staged(r, values, (POLICY, None)) + assert cost.checkpoint_input_gradient == gradient + assert cost.required == int( + (out + retained + gradient + max(decoder, stage) + COLD) * 1.1 + ) + estimate = r._estimate_required_memory_bytes_from_values( + packed_tokens=n, + output_bytes=out, + signature=signature, + logical_tokens=n, + group_rows=groups, + slot_refs=(POLICY, None), + head_workspace_bytes=head, + head_backward_traced=True, + ) + assert estimate == cost.required + # An untraced head keeps the unstaged price as a floor. + unstaged = max(workspace, head) + extra + untraced = priced(r, values, (POLICY, None)) + assert untraced.required == int( + (out + retained + gradient + max(unstaged, decoder, stage) + COLD) * 1.1 + ) + assert untraced.checkpoint_workspace == cost.checkpoint_workspace + # Warm: no first-execution transients and no TE growth in either stage. + r._te_workspace_growth_bytes = lambda: 0 + r._memory_profiles[signature] = _MemoryProfile(bytes_per_token=1, packed_tokens=n) + decoder = r._checkpoint_memory_floor(groups)[1] + extra + assert decoder == workspace - te + extra + stage = head + 2 * gradient + state if head else 0 + assert staged(r, values, (POLICY, None)).required == int( + (out + retained + gradient + max(decoder, stage)) * 1.1 + ) + + +def test_head_bound_short_wave_no_longer_adds_the_adapter_extra(monkeypatch): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + retained, workspace = r._checkpoint_memory_floor(groups) + head = 2 * workspace + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + extra = oracle(pending, [retained // 40] * 40) + assert head > workspace and extra > 800 * MiB + values = (n, out, signature, groups, head) + unstaged = int((out + retained + head + 67 * 2048 * 2 + extra + COLD) * 1.1) + assert staged(r, values, (POLICY, None)).required < unstaged + # Untraced (e.g. a top-k-only head, whose backward keeps its recomputed + # logits beside their gradient): never below the unstaged price. + assert priced(r, values, (POLICY, None)).required == unstaged + + +@pytest.mark.parametrize("uncovered", ["cp4", "uncovered_moe"]) +def test_untraced_recompute_keeps_the_head_in_the_decoder_stage(monkeypatch, uncovered): + r = covered_rank(monkeypatch) + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 4096)) + if uncovered == "cp4": + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 4, 1)) + else: + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 + head = 700 * MiB + retained, workspace = r._checkpoint_memory_floor(groups) + gradient = r._checkpoint_input_gradient_bytes(groups) + assert gradient == retained # One gradient per boundary. + with_pending(monkeypatch, r, ()) + # Even for a traced head. + cost = staged(r, (n, out, signature, groups, head), (POLICY, None)) + assert cost.checkpoint_workspace == max(workspace, head) + COLD + assert cost.required == int( + (out + retained + max(workspace, head) + gradient + COLD) * 1.1 + ) + + +def test_split_keeps_each_head_stage_beside_every_adapter_gradient(monkeypatch): + r = covered_rank(monkeypatch) + head = 700 * MiB + values = r._estimate_flat_forward(requests(67, 4096)) + n, out, signature, groups, _ = values + pending = [23 * MiB] * 40 + [0] + with_pending(monkeypatch, r, pending) + left = staged(r, (n, out, signature, groups, head), (POLICY, None)) + right = staged(r, (n, out, signature, groups, head), (OTHER, None)) + gradient = left.checkpoint_input_gradient + stage = ( + head + + 2 * gradient + + 67 * r._backward_row_state_bytes() + + r._te_workspace_growth_bytes() + ) + assert ( + left.checkpoint_workspace + == max(r._checkpoint_memory_floor(groups)[1], stage) + COLD + ) + # The left child's head can run after the right child's decoder backward: + # both children's boundaries and incoming gradients, the left head stage + # and every adapter gradient the right one allocated (distinct slots). + after_right = ( + left.checkpoint_retained + + right.checkpoint_retained + + 2 * gradient + + stage + + COLD + + right.checkpoint_adapter_gradient + ) + split = TrainerRank._split_required_memory([left, right]) + assert split >= int(after_right * 1.1) + assert split >= int( + (after_right + left.checkpoint_adapter_gradient) * 1.1 + ) # Either may run first, and distinct slots each allocate their own. + + +def test_traced_head_backward_is_target_only_on_the_traced_path(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + target = request(512, grad=True) + # Until the fused statistics have run in this process, nothing is traced. + state = {"succeeded": set(), "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + assert r._head_backward_traced([target], 512) is False + # The top-k kernel's success does not prove the target-only one. + state["succeeded"].add("local_topk_stats") + assert r._head_backward_traced([target], 512) is False + state["succeeded"].add("local_logsumexp_stats") + assert r._head_backward_traced([target], 512) is True + # Top-k, logits and hidden-state outputs keep further dense gradients. + for extra in ({"top_k": 2}, {"logits": True}, {"hidden_states": True}): + other = replace(target, **extra) + assert r._head_backward_traced([target, other], 512) is False + topk_only = replace(target, target_tokens=None, top_k=2) + assert r._head_backward_traced([topk_only], 512) is False + # The fused statistics need 64 rows; bounds that straddle them are open. + assert r._head_backward_traced([target], 63) is False + assert r._head_backward_traced([target], 63, 64) is None + assert r._head_backward_traced([target], 32, 63) is False + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_TOPK", "0") + assert r._head_backward_traced([target], 512) is False + monkeypatch.delenv("ART_TRAINER_RANK_TRITON_TOPK") + # Only CP2 was traced. + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 1, 1)) + assert r._head_backward_traced([target], 512) is False + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # One error sent a kernel to the FP32 fallback silently: never again. + state["failed"] = True + assert r._head_backward_traced([target], 512) is False + state["failed"] = False + r.runtime.model[0].config.use_mup = True + assert r._head_backward_traced([target], 512) is False + + +def test_undecided_heads_bound_both_ways(): + from art.trainer_rank._impl import _traced_states + + assert _traced_states([]) == (False,) + assert _traced_states([True, True]) == (True,) + assert _traced_states([True, False, None]) == (False,) + # Lower bounds take the cheaper, acceptance the dearer of both. + assert _traced_states([True, None]) == (True, False) + + +def test_each_row_keeps_its_rope_embedding_and_index_state(): + import torch + + r = rank() + model = r.runtime.model[0] + assert r._backward_row_state_bytes() == 256 + # Qwen3.6-35B-A3B: a 64-wide rotary embedding (32 frequencies), FP32. + model.rotary_pos_emb = torch.nn.Module() + model.rotary_pos_emb.inv_freq = torch.ones(32) + assert r._backward_row_state_bytes() == 64 * 4 + 256 + + +def test_plan_stages_only_traced_gradient_heads(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + target = request(512, grad=True) + topk_only = replace(target, target_tokens=None, top_k=2) + plans = [r._plan_flat_forward([request]) for request in (target, topk_only)] + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # A top-k-only head's backward keeps its recomputed logits beside their + # gradient, twice the one buffer its workspace prices: never staged. + assert [r._plan_head_backward_traced(plan) for plan in plans] == [True, False] + no_grad = r._plan_flat_forward([request(512)]) + assert r._plan_head_backward_traced(no_grad) is False + + +def test_head_stage_covers_a_small_chunks_fp32_fallback(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + + r = head_rank() + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: True) + r._te_workspace_growth_bytes = lambda: 0 + groups = ((8, True),) + dense = 248320 * 2 + state = 8 * r._backward_row_state_bytes() + # However the CP split leaves a rank's projected rows, a chunk below the + # fused minimum takes the FP32 fallback: nine buffers of up to 63 rows. + stage = r._checkpoint_head_stage_bytes(3 * 8 * dense, 0, groups, None) + assert stage == 9 * 63 * dense + state + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 3 * 512 * dense + state + monkeypatch.setenv("ART_TRAINER_RANK_TRITON_MIN_ROWS", "512") + stage = r._checkpoint_head_stage_bytes(3 * 512 * dense, 0, groups, None) + assert stage == 9 * 511 * dense + state + + +def test_a_staged_plan_runs_its_fused_statistics_strictly(monkeypatch): + from types import SimpleNamespace + + from art.trainer_rank import _impl, topk + + state = {"succeeded": {"local_logsumexp_stats"}, "failed": False} + monkeypatch.setattr(_impl, "_TRITON_STATS_STATE", state) + + def fail(*args, **kwargs): + raise RuntimeError("kernel launch failed") + + monkeypatch.setattr(topk, "local_logsumexp_stats", fail) + chunk: Any = SimpleNamespace(is_cuda=True, shape=(512, 248320)) + # Unstaged: the FP32 fallback, and no later plan stages again. + assert _impl._try_triton_stats("local_logsumexp_stats", chunk) is None + assert state["failed"] is True + # Staged: the plan's price assumed the fused path, so it raises instead. + with pytest.raises(RuntimeError, match="admitted on their memory"): + _impl._try_triton_stats("local_logsumexp_stats", chunk, strict=True) + # Too few rows is a predictable fallback the head stage prices. + small = SimpleNamespace(is_cuda=True, shape=(63, 248320)) + assert _impl._try_triton_stats("local_logsumexp_stats", small, strict=True) is None + + +def test_execution_binds_strictness_to_the_staged_admission(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import ForwardOutput, _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + monkeypatch.setattr(r, "_moe_recompute_covered_for", lambda ref: True) + plan = r._plan_flat_forward([request(512, grad=True), request(16, hidden=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + monkeypatch.setattr(r, "_topology", lambda: None) + monkeypatch.setattr(r, "_validate_hybridep_topology", lambda: None) + monkeypatch.setattr(r, "_configure_hybridep", lambda *a, **k: None) + monkeypatch.setattr(r, "_prepare_packed_forward", lambda packed: None) + seen = [] + + def forward_packed(items, prepared): + seen.append(_impl._HEAD_STATISTICS_STRICT.get()) + return [ForwardOutput(None, None, None, None)] * len(items) + + monkeypatch.setattr(r, "_forward_packed", forward_packed) + r._execute_flat_plan(plan) + assert seen == [False, False] # Never priced staged. + assert r._plan_head_backward_traced(plan) is True + assert plan._head_staged is True + seen.clear() + r._execute_flat_plan(plan) + # Only the gradient group runs strictly; nothing leaks past execution. + grad_first = [group.grad_enabled for group in plan.groups] + assert seen == grad_first + assert _impl._HEAD_STATISTICS_STRICT.get() is False + # A later price that no longer stages (a kernel failed since) cannot + # weaken the admission that relied on it. + _impl._TRITON_STATS_STATE["failed"] = True + assert r._plan_head_backward_traced(plan) is False + assert plan._head_staged is True + + +def test_an_eligible_head_the_price_does_not_stage_is_not_strict(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + plan = r._plan_flat_forward([request(512, grad=True)]) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + # The recompute is not covered: the head keeps the unstaged price. + monkeypatch.setattr(r, "_checkpoint_gradient_covered", lambda *a, **k: False) + assert r._plan_head_backward_traced(plan) is True + assert getattr(plan, "_head_staged", False) is False + + +def test_several_labels_per_row_are_not_traced(monkeypatch): + from test_trainer_rank_head_memory import rank as head_rank + from test_trainer_rank_head_memory import request + + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = head_rank() + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 2, 1)) + single = request(512, grad=True) + several = replace( + single, target_tokens=single.target_tokens.unsqueeze(1).repeat(1, 4) + ) + assert r._head_backward_traced([single], 512) is True + assert r._head_backward_traced([several], 512) is False + # One label per token over a leading batch axis is still one per row. + batched = replace(single, input_tokens=single.input_tokens.unsqueeze(0)) + assert r._head_backward_traced([batched], 512) is True diff --git a/tests/unit/test_trainer_rank_moe_memory.py b/tests/unit/test_trainer_rank_moe_memory.py index 3fbdcee89..7bc27d8fd 100644 --- a/tests/unit/test_trainer_rank_moe_memory.py +++ b/tests/unit/test_trainer_rank_moe_memory.py @@ -1,6 +1,7 @@ """CPU contracts for one known MoE component, not a whole-model memory bound.""" from dataclasses import replace +import math from types import SimpleNamespace from typing import Any, cast import weakref @@ -9,8 +10,12 @@ import torch from art.trainer_rank import ForwardInput, TrainerRank +from art.trainer_rank._impl import ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) from art.trainer_rank._impl import ( _PACKED_PRICED_LOGICAL_ROW_BYTES, + _ep_routed_row_allowance, _MemoryProfile, _MemorySignature, _moe_output_bytes_per_token, @@ -203,20 +208,33 @@ def _hybridep(layer, ep: int, manager: str = "hybridep"): return layer -@pytest.mark.parametrize("ep", [2, 4]) -def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep): +@pytest.mark.parametrize("ep,allowance", [(2, 1.4), (4, 1.6), (8, 2.0), (16, 2.4)]) +def test_hybridep_prices_routed_rows_with_imbalance_allowance(layer, ep, allowance): # HybridEP gives each rank the pairs routed to its local experts: balanced - # routing matches EP1's top-k rows per local token, with a 1.5x allowance. + # routing matches EP1's top-k rows per local token, scaled by the measured + # EP-dependent worst-layer imbalance (log2 growth beyond EP8). single = _moe_output_bytes_per_token([layer], ParallelShape(tp=1, cp=1)) sharded = ParallelShape(tp=1, cp=ep, ep=ep) expert = _moe_output_bytes_per_token([_hybridep(layer, ep)], sharded) - # Top-k 8 routed rows per local token become 12; no shared experts here. - assert single > 0 and expert == single // 8 * 12 + # Top-k 8 routed rows per local token become 8 x allowance; no shared + # experts here. + assert _ep_routed_row_allowance(ep) == pytest.approx(allowance) + assert single > 0 and expert == math.ceil(single * allowance) + + +def test_unmeasured_ep_sizes_use_the_next_measured_allowance(): + assert _ep_routed_row_allowance(1) == 1.0 + assert _ep_routed_row_allowance(3) == _ep_routed_row_allowance(4) == 1.6 + assert _ep_routed_row_allowance(6) == 2.0 + # Never above every pair on one rank. + assert _ep_routed_row_allowance(1024) <= 1024 def test_hybridep_keeps_the_enclosing_fc1_stage(layer): - # The FC1 inputs and gate/up sum stay live at the FC2 sum under HybridEP - # too; EP1's two dispatched H-wide inputs over-count HybridEP's one. + # The FC1 input and gate/up sum stay live at the FC2 sum under HybridEP + # too. The EP1 all-to-all path holds two routed H-wide inputs (its + # permuted rows and their expert-sorted copy); HybridEP permutes while + # dispatching and holds one, as Qwen3.6 CP2 allocator traces show. single = _moe_output_bytes_per_token( [_enclosing_moe(layer)], ParallelShape(tp=1, cp=1) ) @@ -224,7 +242,7 @@ def test_hybridep_keeps_the_enclosing_fc1_stage(layer): expert.token_dispatcher.num_local_experts = 128 sharded = _moe_output_bytes_per_token([expert], ParallelShape(tp=1, cp=2, ep=2)) assert single == 8 * (512 + 3 * 2048 + 2 * 2048 + 1024) * 2 - assert sharded == single // 8 * 12 + assert sharded == math.ceil(8 * 1.4 * (512 + 3 * 2048 + 2048 + 1024) * 2) @pytest.mark.parametrize( @@ -591,7 +609,7 @@ def held(rows): monkeypatch.setattr( rank, "_checkpoint_memory_floor", - lambda rows, refs=None, segments=0: (10**7, 10**6), + lambda rows, refs=None, segments=0, routed_rows=None: (10**7, 10**6), ) grown = rank._plan_cost(plan) monkeypatch.setattr(rank, "_plan_hybridep_growth_bytes", lambda plan: 0) @@ -610,6 +628,92 @@ def held(rows): assert rank._plan_hybridep_growth_bytes(plan) == 0 +def test_converted_stages_follow_routed_rows_and_shared_stays_local(): + rank = _rank() + rank._moe_output_bytes_per_token = 1000 + rank._moe_forward_stages = ((1200, 50),) + rank._moe_forward_shared_bytes = 100 + assert rank._moe_workspace_bytes(10) == 10 * 1200 + 50 + assert rank._moe_workspace_bytes(10, routed_rows=6) == 6 * 1200 + 50 + 4 * 100 + assert rank._moe_workspace_bytes(10, routed_rows=0) == 50 + 10 * 100 + + +def test_ep_group_must_be_exactly_the_cp_group(monkeypatch): + ps = pytest.importorskip("megatron.core.parallel_state") + from art.trainer_rank import _impl + + shape = ParallelShape(tp=1, cp=2, ep=2) + for other in ( + replace(shape, ep=1), + replace(shape, ep=4), + replace(shape, tp=2), + replace(shape, etp=2), + ): + assert not _impl._ep_group_is_cp_group(other) + # Without initialized process groups there is nothing to compare. + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(_impl.dist, "is_initialized", lambda: True) + expert, context = object(), object() + ranks = {id(expert): [0, 1], id(context): [1, 0]} + monkeypatch.setattr(ps, "get_expert_model_parallel_group", lambda **_: expert) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: context) + monkeypatch.setattr(_impl.dist, "get_process_group_ranks", lambda g: ranks[id(g)]) + assert _impl._ep_group_is_cp_group(shape) + # EP spanning ranks outside this CP group sees other batches' rows. + ranks[id(context)] = [0, 2] + assert not _impl._ep_group_is_cp_group(shape) + monkeypatch.setattr(ps, "get_context_parallel_group", lambda **_: None) + assert not _impl._ep_group_is_cp_group(shape) + + +def test_routed_rows_use_the_cp_group_share_only_when_it_is_the_ep_group( + monkeypatch, +): + rank = _rank() + plan = rank._plan_flat_forward( + [ForwardInput(input_tokens=torch.arange(64), target_tokens=torch.arange(64))] + ) + assert rank._plan_group_routed_rows(plan) == (64,) + plan = replace(plan, signature=replace(plan.signature, topology=(1, 1, 2, 1))) + monkeypatch.setattr(rank, "_topology", lambda: SimpleNamespace(tp=1, cp=2)) + monkeypatch.setattr(rank, "_plan_group_rows", lambda plan: ((52480, True),)) + totals = [96794] + monkeypatch.setattr( + rank, "_cp_group_model_tokens", lambda batch, topology: totals[0] + ) + assert rank._plan_group_routed_rows(plan) == (52480,) + rank._ep_group_is_cp_group = True + assert rank._plan_group_routed_rows(plan) == (48397,) + totals[0] = 10**6 + assert rank._plan_group_routed_rows(plan) == (52480,) + + +def test_cp_group_total_is_its_larger_layout(monkeypatch): + runtime = pytest.importorskip("art.megatron.context_parallel.runtime") + bundle = SimpleNamespace( + token_layout_index=SimpleNamespace(token_counts_by_rank=(52480, 44314)) + ) + monkeypatch.setattr( + runtime, + "_get_or_build_planning_bundle", + lambda **_: ("key", bundle, None, None), + ) + monkeypatch.setattr( + runtime, + "_plan_gdn_global_execution", + lambda **_: SimpleNamespace(gdn_token_counts_by_rank=(48100, 48800)), + ) + values: dict[str, Any] = dict( + group_ids=None, parent_ids=None, topology=None, config=None, original_seq_len=0 + ) + total = runtime.context_parallel_model_token_total + assert total(**values, build_gdn_execution_spec=False) == 96794 + assert total(**values, build_gdn_execution_spec=True) == 96900 + # An empty rank still dispatches, and routes, one padding row. + bundle.token_layout_index.token_counts_by_rank = (2, 0) + assert total(**values, build_gdn_execution_spec=False) == 3 + + def test_split_charges_the_largest_hybridep_growth_beside_any_child_peak(): from art.trainer_rank._impl import _SubforwardCost @@ -697,7 +801,8 @@ def hybrid_checkpoint_rank(layer, monkeypatch): ), ) ) - assert rank._moe_output_bytes_per_token == 282624 + # This branch prices HybridEP routed rows on the EP group's balanced share. + assert rank._moe_output_bytes_per_token == 217908 assert rank._moe_memory_supported return rank @@ -739,17 +844,22 @@ def test_hybridep_recompute_prices_fresh_dense_output_without_buffer_growth( ) assert rank._plan_hybridep_growth_bytes(plan) == 0 retained, workspace = rank._checkpoint_memory_floor(groups) - assert workspace == 218752 * 2048 * 2 == 896008192 + # The combine output, with the TE workspaces live beside it. + assert workspace == 218752 * 2048 * 2 + rank._te_workspace_growth_bytes() cost = rank._subforward_cost(**values) - assert cost.required == int((8 + 2 * retained + workspace) * 1.1) - assert cost.checkpoint_workspace == workspace # Maximum, not stage + output. + # Unprofiled: the first execution's transients sit beside the extent. + assert cost.required == int((8 + 2 * retained + workspace + COLD) * 1.1) + assert cost.checkpoint_workspace == workspace + COLD # Not stage + output. rank._available_memory_bytes = lambda: 600000000 assert rank._memory_check_required(baseline.required).fits assert not rank._memory_check_required(cost.required).fits rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=1, packed_tokens=2 ) - assert rank._subforward_cost(**values).required == cost.required + # Profiled: the floor alone, without first-execution transients. + assert rank._subforward_cost(**values).required == int( + (8 + 2 * retained + workspace) * 1.1 + ) rank._memory_profiles[signature] = _MemoryProfile( bytes_per_token=10**9, packed_tokens=2 ) diff --git a/tests/unit/test_trainer_rank_pending_memory.py b/tests/unit/test_trainer_rank_pending_memory.py index 359de8daa..35f70c535 100644 --- a/tests/unit/test_trainer_rank_pending_memory.py +++ b/tests/unit/test_trainer_rank_pending_memory.py @@ -12,6 +12,7 @@ from art.megatron.prefix_tree_packing import prefix_tree_pack from art.trainer_rank import ForwardInput, TrainerRank from art.trainer_rank import _gdn_memory as g +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import Unset, _MemoryProfile @@ -21,7 +22,7 @@ def module(cls): return obj -def rank_with_moe(moe_layer, *, install_hooks=False): +def rank_with_moe(moe_layer, *, install_hooks=False, stand_in=True): from megatron.core.ssm.gated_delta_net import GatedDeltaNet from megatron.core.transformer.transformer_block import TransformerBlock from transformer_engine.pytorch import RMSNorm @@ -96,6 +97,10 @@ def rank_with_moe(moe_layer, *, install_hooks=False): ) ) r._dp_rank_and_size = lambda: (0, 1) # Uninitialized MCore has no CPU DP group. + if stand_in: + # The one MoE layer stands in for all 40 of Qwen3.6-35B-A3B's: recompute + # is covered if that layer prices its FC1 stage too. + r._moe_recompute_covered = r._moe_gradient_enclosed == (True,) return r, gd @@ -133,7 +138,7 @@ def test_actual_constructor_cache_and_full_plan(pending_rank): assert ( rank._memory_check(plan).estimated_required_bytes == rank._plan_cost(plan).required - == 32229502659 + == 23404942633 ) selected = rank._select_next_micro_batch(requests, 0) assert ( @@ -196,7 +201,7 @@ def test_exact_pending_demand_survives_recovery(monkeypatch, pending_rank, fits_ plan = pending_rank._plan_flat_forward(requests) assert pending_rank._estimate_flat_forward(requests) is None assert g.plan_floor(pending_rank, plan) == (8296857600, 12705630112) - assert pending_rank._memory_check(plan).estimated_required_bytes == 32229502659 + assert pending_rank._memory_check(plan).estimated_required_bytes == 23404942633 _check_component_demand_recovery( monkeypatch, pending_rank, requests, fits_after=fits_after ) @@ -214,8 +219,8 @@ def test_original_installed_norm_preserves_pending_floor(layer): assert g.model_shapes(rank) is not None plan = rank._plan_flat_forward(full_requests()) assert g.plan_floor(rank, plan) == (8296857600, 12705630112) - assert rank._memory_check(plan).estimated_required_bytes == 32229502659 - assert rank._plan_cost(plan).required == 32229502659 + assert rank._memory_check(plan).estimated_required_bytes == 23404942633 + assert rank._plan_cost(plan).required == 23404942633 assert rank._estimate_flat_forward(full_requests()) is None for requests in ([], full_requests(no_grad=True)): assert g.plan_floor(rank, rank._plan_flat_forward(requests)) == (0, 0) @@ -382,9 +387,11 @@ def test_constructor_declined_moe_keeps_generic_admission(layer, unsupported): plan = rank._plan_flat_forward(requests) assert g.plan_floor(rank, plan) == (0, 0) required = rank._plan_cost(plan).required - # Generic checkpoint-input accounting still applies without a MoE component. + # Generic checkpoint accounting, including the recomputed layer's mixer, + # still applies without a MoE component. gradient = 50640 * 40 * 2048 * 2 - assert required == int((plan.output_bytes + 2 * gradient) * 1.1) + mixer = 50640 * rank._recomputed_mixer_bytes_per_token() + assert required == int((plan.output_bytes + 2 * gradient + mixer + COLD) * 1.1) rank._available_memory_bytes = lambda: required - 1 assert not rank._memory_check(plan).fits rank._available_memory_bytes = lambda: required diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index a58adfbc1..a9b26b6af 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -299,6 +299,8 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_input_gradient": 0, "checkpoint_peak_increment": 0, "hybridep_growth": 0, + "checkpoint_adapter_gradient": 0, + "checkpoint_adapter_gradient_slots": "", }, } ], @@ -387,6 +389,7 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "checkpoint_workspace", "checkpoint_peak_increment", "hybridep_growth", + "checkpoint_adapter_gradient", ): altered = json.loads(path.read_bytes()) altered["replay"]["memory_replay"]["estimates"][0]["cost_components"][ diff --git a/tests/unit/test_trainer_rank_shared_memory.py b/tests/unit/test_trainer_rank_shared_memory.py index 36bc971d7..9016f6cc8 100644 --- a/tests/unit/test_trainer_rank_shared_memory.py +++ b/tests/unit/test_trainer_rank_shared_memory.py @@ -1,5 +1,6 @@ """One supported shared return held across routed compute; not all backward saves.""" +import math from types import SimpleNamespace import pytest @@ -106,6 +107,9 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): assert rank._moe_output_bytes_per_token == 192512 checkpoint_coefficient = 196608 if gate else 192512 assert rank._moe_checkpoint_grad_bytes_per_token == checkpoint_coefficient + # The shared return is the part that stays on local rows under HybridEP. + assert rank._moe_forward_shared_bytes == 4096 + assert rank._moe_gradient_shared_bytes == (8192 if gate else 4096) shapes = g.model_shapes(rank) assert shapes is not None and shapes[1][0].moe_bytes_per_row == 192512 requests = full_requests(no_grad) @@ -126,7 +130,8 @@ def test_shared_return_in_actual_constructor_and_plan(layer, gate, no_grad): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - assert rank._plan_cost(plan).required == (32685829827 if gate else 32457666243) + # The incoming gradient replaces one gradient per boundary. + assert rank._plan_cost(plan).required == (23861269801 if gate else 23633106217) selected = rank._select_next_micro_batch(requests, 0) assert ( selected.check.estimated_required_bytes @@ -147,7 +152,7 @@ def test_original_norm_installation_preserves_shared_return(layer, gated): 8296857600, 50640 * (checkpoint_coefficient + 128) + 3157761952, ) - expected = 32685829827 if gated else 32457666243 + expected = 23861269801 if gated else 23633106217 assert rank._memory_check(plan).estimated_required_bytes == expected assert rank._plan_cost(plan).required == expected @@ -305,14 +310,18 @@ def test_pre_gate_cache_precedes_owned_dispatcher_and_is_checkpoint_only(layer): ) # Installed dispatcher partials must not be repriced. assert rank._moe_checkpoint_grad_bytes_per_token == 196608 groups = ((19, True), (23, False)) + # Gradient rows also keep the recomputed layer's residual and norm rows and + # its MoE routing state; the first call allocates TE's cuBLAS workspaces. + beside = 2 * 2048 * 2 + rank._moe_checkpoint_state_bytes_per_token() assert rank._checkpoint_memory_floor(groups) == ( 19 * 40 * 4096, max( - 19 * (196608 + 128), + 19 * (196608 + 128 + beside), 23 * (192512 + 4 * 2048 * 2), - 19 * (196608 - 32768 + 128) + 10485760, + 19 * (196608 - 32768 + 128 + beside) + 10485760, 23 * (192512 - 32768 + 128 + 4 * 2048 * 2) + 10485760, - ), + ) + + rank._te_workspace_growth_bytes(), ) for mode in (None, "selective"): rank.runtime.model[0].decoder.config.recompute_granularity = mode @@ -340,7 +349,7 @@ def test_pre_gate_mixed_reference_and_exact_cost_mode_selection(layer, gradient_ ) assert rank._checkpoint_memory_floor(rank._plan_group_rows(mixed)) == ( 67 * 40 * 4096, - 4096 * (192512 + 4 * 2048 * 2), + 4096 * (192512 + 4 * 2048 * 2) + rank._te_workspace_growth_bytes(), ) # A reference-only path must not read or validate the unused gradient cache. rank._moe_checkpoint_grad_bytes_per_token = None @@ -410,11 +419,17 @@ def test_shared_return_escapes_the_ep_routed_allowance(layer, checkpoint_grad, s layer.config.context_parallel_size = layer.config.expert_model_parallel_size = 2 # EP2 halves the experts each rank owns. _hybridep(layer, 2).token_dispatcher.num_local_experts = local_experts // 2 - # HybridEP's 1.5x allowance turns top-k 8 into 12 routed rows; the gated - # shared return (doubled for checkpoint backward) is per local token. + # HybridEP's EP2 allowance (1.4) turns top-k 8 into 11.2 routed rows, each + # with one dispatched H-wide input; the gated shared return (doubled for + # checkpoint backward) is per local token. + collected: list[int] = [] assert ( _moe_output_bytes_per_token( - [layer], ParallelShape(tp=1, cp=2, ep=2), checkpoint_grad=checkpoint_grad + [layer], + ParallelShape(tp=1, cp=2, ep=2), + checkpoint_grad=checkpoint_grad, + shared_bytes=collected, ) - == 12 * 11776 * 2 + shared + == math.ceil(8 * 1.4 * 9728 * 2) + shared ) + assert collected == [shared] diff --git a/tests/unit/test_trainer_rank_slot_memory.py b/tests/unit/test_trainer_rank_slot_memory.py index 71ed04224..04bf8c06e 100644 --- a/tests/unit/test_trainer_rank_slot_memory.py +++ b/tests/unit/test_trainer_rank_slot_memory.py @@ -136,19 +136,27 @@ def test_subforward_cost_reuses_floor_without_caching_across_slot_changes( original = rank._checkpoint_memory_floor calls = [] - def floor(*args): - result = original(*args) + def floor(*args, **kwargs): + result = original(*args, **kwargs) calls.append(result) return result monkeypatch.setattr(rank, "_checkpoint_memory_floor", floor) + rows = rank._plan_group_rows(plan) + refs = tuple(group.slot_ref for group in plan.groups) + + def gradient(retained): + # Derived from the one floor call's boundaries, without another call. + return rank._checkpoint_input_gradient_bytes(rows, refs, retained=retained) + before = rank._plan_cost(plan) assert len(calls) == 1 - assert before.checkpoint_input_gradient == calls[0][0] + assert before.checkpoint_input_gradient == gradient(calls[0][0]) > 0 load_slot(rank, "selected", 64) after = rank._plan_cost(plan) assert len(calls) == 2 - assert after.checkpoint_input_gradient == calls[1][0] + assert after.checkpoint_input_gradient == gradient(calls[1][0]) > 0 + assert len(calls) == 2 assert calls[1][1] > calls[0][1] assert after.checkpoint_workspace > before.checkpoint_workspace assert after.required > before.required @@ -273,3 +281,19 @@ def guarded(name, *args, **kwargs): monkeypatch.setattr(builtins, "__import__", guarded) assert rank._slot_memory_shapes(ref) == () + + +def test_partial_slot_with_unpriced_fc1_stages_keeps_boundary_gradients(layer): + # A slot with FC1 adapters but no FC2 adapter still prices FC2 rows from the + # original metadata, but its FC1 converted weights go unpriced: keep one + # gradient per boundary for it. + rank, _ = rank_with_moe(weights(layer, 8)) + ref = load_slot(rank, "partial", 8) + groups = ((100, True),) + assert rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 2048 * 2 + del layer.experts.linear_fc2.lora._slot_keys[ref] + assert rank._moe_workspace_bytes(1, checkpoint_grad=True, slot_ref=ref) > 0 + assert not rank._moe_recompute_covered_for(ref) + assert rank._checkpoint_input_gradient_bytes(groups, (ref,)) == 100 * 40 * 4096 + assert rank._checkpoint_input_gradient_bytes(groups) == 100 * 2048 * 2 diff --git a/tests/unit/test_trainer_rank_tp_floor.py b/tests/unit/test_trainer_rank_tp_floor.py index 6d5f4d00d..c3f49e3b9 100644 --- a/tests/unit/test_trainer_rank_tp_floor.py +++ b/tests/unit/test_trainer_rank_tp_floor.py @@ -14,6 +14,7 @@ import torch from art.trainer_rank import TrainerRank +from art.trainer_rank._impl import _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD from art.trainer_rank._impl import _MemorySignature H, F, LAYERS = 5120, 17408, 64 @@ -114,10 +115,11 @@ def test_the_traced_tp4_wave_prices_its_boundary_shards_and_their_repeat(): assert workspace == 3 * SEGMENT cost = _required(r) assert cost.checkpoint_input_gradient == retained - # One segment plus up to three TP-padding roots, each with its states. - state = cost.checkpoint_workspace + # One segment plus up to three TP-padding roots, each with its states, + # beside the unprofiled first execution's transients. + state = cost.checkpoint_workspace - COLD assert state == 4 * SEGMENT - assert cost.required == int((OUTPUT + 2 * retained + state) * 1.1) + assert cost.required == int((OUTPUT + 2 * retained + state + COLD) * 1.1) # Measured cold on all four ranks: 7.130 GB (7.060 GB in production), all # but the boundaries a transient recompute workspace; this raw floor # (8.43 GB) covers it. Today's cold admission was 4.637 GB. @@ -216,9 +218,9 @@ def test_gdn_segment_states_are_priced_with_the_segments(): r = tp_rank() rows = 8192 cost = _required(r, group_rows=((rows, True),), gdn_segments=4096) - assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + assert cost.checkpoint_workspace == (4096 + 3) * SEGMENT + COLD assert cost.required == int( - (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT) * 1.1 + (OUTPUT + 2 * rows // 4 * LAYERS * H * 2 + (4096 + 3) * SEGMENT + COLD) * 1.1 ) @@ -227,7 +229,7 @@ def test_tp_padding_roots_carry_their_own_states(): r = tp_rank() cost = _required(r, group_rows=((4, True),), gdn_segments=1) # Four roots' initial states alone: 4 x 12 value heads x 128 x 128 x fp32. - assert cost.checkpoint_workspace == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 + assert cost.checkpoint_workspace - COLD == 4 * SEGMENT > 4 * 12 * 128 * 128 * 4 assert cost.required > 4 * SEGMENT