diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index 9005b367c..4be496758 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -243,6 +243,7 @@ jobs: tests/unit/test_trainer_rank_converted_memory.py \ tests/unit/test_trainer_rank_layout_memory.py \ tests/unit/test_context_parallel_retained_bytes.py \ + tests/unit/test_trainer_rank_dense_memory.py \ tests/unit/test_trainer_rank_split.py \ tests/unit/test_megatron_compile_garbage.py \ tests/unit/test_trainer_rank_cache_recovery.py::test_dense_cp_exact_demand_fits_after_recovery \ @@ -291,4 +292,5 @@ jobs: --ignore=tests/unit/test_megatron_compile_garbage.py \ --ignore=tests/unit/test_trainer_rank_converted_memory.py \ --ignore=tests/unit/test_trainer_rank_layout_memory.py \ - --ignore=tests/unit/test_context_parallel_retained_bytes.py + --ignore=tests/unit/test_context_parallel_retained_bytes.py \ + --ignore=tests/unit/test_trainer_rank_dense_memory.py diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index af1b5bdfd..29130b264 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1640,6 +1640,220 @@ def _hybridep_buffer_bytes(capacity: int, ranks: int, hidden: int, experts: int) return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) +# Largest LoRA rank the dense stage prices (rank-wide intermediates included). +_DENSE_LORA_RANK_LIMIT = 256 + + +def _dense_mlp_recompute_bytes_per_token( + model: Sequence[torch.nn.Module], + slot_ref: "LoRASlotRef | None" = None, + *, + hidden_size: int | None = None, +) -> tuple[int, int]: + """Per-row dense MLP bytes: (gradient recompute stage, no-grad transient). + + Qwen3.8-27B CP2 allocator traces (dense, gated SwiGLU, LoRA on FC1 and FC2): + + - A recomputed layer's peak sits in its FC1 stage: the base output, the + LoRA gate/up output and their sum (2F each) plus one F-wide tensor, 7F + per row (238 KB measured). Early in real q062 runs one rank at a time + held about one more FC1 triplet (6F; +4.8 GB at 23,552 rows), as a + recompile can leave a graph's outputs live; that is priced too. + - A no-grad layer holds its three 2F FC1 tensors, residual, norm and CP + gather rows: 263 KB per row measured, priced as 6F + 6H. + + Both add the rank intermediates of every adapter in the layer. Every + decoder layer, all it runs and ``slot_ref``'s adapters must match the + traced execution (ART's own GDN layer, mixer and norm wrappers included); + otherwise (0, 0) keeps today's allowances. + """ + if len(model) != 1: + return 0, 0 + try: + decoder = _language_model(model[0]).decoder + except (AttributeError, RuntimeError): + return 0, 0 + layers = getattr(decoder, "layers", None) + if not layers or not all(hasattr(layer, "mlp") for layer in layers): + return 0, 0 + try: + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TERowParallelLinear, + ) + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + from megatron.core.transformer.attention import SelfAttention + from megatron.core.transformer.mlp import MLP + from megatron.core.transformer.transformer_block import TransformerBlock + from megatron.core.transformer.transformer_layer import TransformerLayer + + from art.megatron.gdn.operator import ( + _empty_safe_norm_forward, + _gdn_island_layer_forward, + _prefix_tree_forward, + ) + from art.megatron.lora import ( + LoRA, + SelfAttentionLinearProjLoRA, + SharedExpertsLinearFC1LoRA, + SharedExpertsLinearFC2LoRA, + ) + except ImportError: + # Without the traced owner types nothing can match; keep the allowance. + return 0, 0 + + def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: + """No hooks, and no forward but the class's or ART's traced wrapper, + which must still call the class's own forward.""" + forward = vars(module).get("forward") + if module._forward_hooks or module._forward_pre_hooks: + return False + if forward is None: + return True + inner = vars(module).get(delegate) + # Training compile replaces the delegate with Dynamo's wrapper (the + # traced run was compiled); judge the callable it wraps. + while hasattr(inner, "_torchdynamo_orig_callable"): + inner = inner._torchdynamo_orig_callable + return ( + wrapper is not None + and type(forward) is MethodType + and forward.__self__ is module + and forward.__func__ is wrapper + and type(inner) is MethodType + and inner.__self__ is module + and inner.__func__ is type(module).forward + ) + + # The traced hybrid's exact mixer types (Qwen3.5-family attention is + # Megatron Bridge's Qwen3VLSelfAttention); subclasses are unmeasured. + mixers: set[type] = {SelfAttention, GatedDeltaNet} + try: + from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.attention import ( + Qwen3VLSelfAttention, + ) + except ImportError: + pass + else: + mixers.add(Qwen3VLSelfAttention) + + if type(decoder) is not TransformerBlock or not plain(decoder): + return 0, 0 + expected = { + "gated_linear_unit": True, + "params_dtype": torch.bfloat16, + "add_bias_linear": False, + "sequence_parallel": False, + # Fused SwiGLU (bias_swiglu_impl), as Megatron Bridge's Qwen3.5 + # providers configure it and the traced run executed. + "bias_activation_fusion": True, + "use_te_activation_func": False, + "cpu_offloading": False, + "cuda_graph_impl": "none", + "tensor_model_parallel_size": 1, + "pipeline_model_parallel_size": 1, + } + width = hidden = rank = 0 + for layer in layers: + mlp = getattr(layer, "mlp", None) + config = getattr(mlp, "config", None) + fc1, fc2 = getattr(mlp, "linear_fc1", None), getattr(mlp, "linear_fc2", None) + row = getattr(fc2, "row_parallel_lora", None) + adapters = ( + getattr(fc1, "gate_lora", None), + getattr(fc1, "up_lora", None), + getattr(row, "lora", None), + ) + sites = ( + (mlp, MLP), + (fc1, SharedExpertsLinearFC1LoRA), + (getattr(fc1, "linear_fc1", None), TELayerNormColumnParallelLinear), + (fc2, SharedExpertsLinearFC2LoRA), + (row, SelfAttentionLinearProjLoRA), + (getattr(row, "linear_proj", None), TERowParallelLinear), + *((adapter, LoRA) for adapter in adapters), + ) + ffn = getattr(config, "ffn_hidden_size", None) + size = getattr(config, "hidden_size", None) + mixer = getattr(layer, "self_attention", None) + if ( + type(layer) is not TransformerLayer + or not plain( + layer, _gdn_island_layer_forward, "_art_gdn_island_physical_forward" + ) + or type(mixer) not in mixers + or not plain(mixer, _prefix_tree_forward, "_art_physical_forward") + or any(type(site) is not cls for site, cls in sites) + or not all(plain(site) for site, _ in sites) + or type(ffn) is not int + or ffn <= 0 + or type(size) is not int + or size <= 0 + or (hidden_size is not None and size != hidden_size) + or getattr(fc1, "non_gated", None) is not False + or getattr(fc1, "out_features", None) != 2 * ffn + or any( + type(getattr(config, name, None)) is not type(value) + or getattr(config, name) != value + for name, value in expected.items() + ) + or getattr(config, "fp8", None) + or getattr(config, "fp4", None) + or getattr(config, "activation_func", None) is not torch.nn.functional.silu + or getattr(mlp, "activation_func", None) is not torch.nn.functional.silu + or getattr(config, "activation_func_clamp_value", None) is not None + or getattr(config, "glu_linear_offset", 0.0) != 0.0 + ): + return 0, 0 + # Everything else the layer runs, the mixer's children included, must + # be the traced execution too: no hooks, and no forward but ART's + # empty-safe norm wrapper. Every adapter must be an exact LoRA whose + # selector is the one execution uses, within the priced rank. + layer_rank = 0 + for child in layer.modules(): + if child is layer or child is mixer: + continue + if not ( + plain(child) + or plain( + child, + _empty_safe_norm_forward, + "_art_empty_safe_norm_physical_forward", + ) + ): + return 0, 0 + if not isinstance(child, LoRA): + continue + if type(child) is not LoRA or any( + name in vars(child) for name in ("_slot", "active_lora_tensors") + ): + return 0, 0 + tensors = _slot_lora_tensors(child, slot_ref) + if tensors is None: + if slot_ref is None or slot_ref.name is None: + return 0, 0 + continue # This slot has no adapter here: base output only. + a, b = tensors + if ( + not isinstance(a, torch.Tensor) + or not isinstance(b, torch.Tensor) + or a.ndim != 2 + or b.ndim != 2 + or a.shape[1] != b.shape[0] + or not 0 < a.shape[1] <= _DENSE_LORA_RANK_LIMIT + ): + return 0, 0 + layer_rank += int(a.shape[1]) + rank = max(rank, layer_rank) + width, hidden = max(width, ffn), max(hidden, size) + # Each of a layer's adapters, the mixer's too, keeps its rank-wide input + # product and gradient. + adapters = 2 * rank + return (7 * width + 6 * width + adapters) * 2, ( + 6 * width + 6 * hidden + adapters + ) * 2 + + def _moe_output_bytes_per_token( model: Sequence[torch.nn.Module], shape: ParallelShape, @@ -2099,6 +2313,18 @@ def memory_field(name: str, default: Any = None) -> Any: self._moe_recompute_covered = len( self._moe_gradient_enclosed ) == self._num_layers and all(self._moe_gradient_enclosed) + # Dense models whose every layer is the traced gated MLP price that + # stage instead, and with it one input gradient. + ( + self._dense_recompute_bytes_per_token, + self._dense_no_grad_bytes_per_token, + ) = ( + (0, 0) + if self._moe_layers + else _dense_mlp_recompute_bytes_per_token( + runtime.model, hidden_size=self._hidden_size + ) + ) self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, @@ -3760,6 +3986,7 @@ def _split_chunk_lower_cost( # The average CP load is an optimistic bound, not an admission cost. retained_tokens=(packed_tokens + signature.topology[2] - 1) // signature.topology[2], + lower_bound=True, ) profile = self._memory_profiles.get(signature) if ( @@ -4275,51 +4502,32 @@ def _sequence_parallel_floor_covered(self, layers: int, tp: int, cp: int) -> boo workspace = 2 * hidden * tp + stage + max(attention, gdn) + 2 * hidden * tp return layers * hidden >= workspace + hidden - def _checkpoint_memory_floor( - self, - group_rows: tuple[tuple[int, bool], ...], - slot_refs: tuple["LoRASlotRef | None", ...] | None = None, - gdn_segments: int = 0, - routed_rows: tuple[int, ...] | None = None, - layouts: tuple[_GroupLayout, ...] | None = None, - ) -> tuple[int, int]: - """Conservative saved-boundary charge and one recomputed layer's workspace. + def _checkpoint_floor_decoder( + self, *, sequence_parallel: bool = False + ) -> Any | None: + """The decoder whose saved boundaries the checkpoint floor prices, or None. - ``routed_rows`` are each group's balanced dispatched rows per rank - (``_plan_group_routed_rows``); by default, its local rows. With - ``layouts`` (``_plan_group_layouts``), price every rank on its own CP - layouts instead of the busiest rank's rows (``_layout_checkpoint_floor``). - Count actual local full/uniform/1 boundaries, including aliases, rather - than claiming measured distinct storage. Only this call's new groups - enter the term; already-live graphs remain in the availability baseline. - 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). + Every local layer recomputed full/uniform/1 in BF16 at PP1, with no + custom checkpointed forward: at TP1 by default, which the dense widths + also require, so a discount never outlives the floor that carries its + TE growth; or, with ``sequence_parallel``, at a TP > 1 that + ``_sequence_parallel_floor_covered`` holds for. """ - gradient_rows = sum(rows for rows, grad in group_rows if grad) - if not group_rows or len(self.runtime.model) != 1: - return 0, 0 + if len(self.runtime.model) != 1: + return None try: decoder = _language_model(self.runtime.model[0]).decoder except (AttributeError, RuntimeError): - return 0, 0 + return None try: from megatron.core.transformer.transformer_block import TransformerBlock except ModuleNotFoundError as error: if error.name != "megatron": raise - return 0, 0 + return None if type(decoder) is not TransformerBlock: - return 0, 0 + return None config = decoder.config layers = len(decoder.layers) _, tp, cp, pp = self._topology_key() @@ -4328,7 +4536,7 @@ def _checkpoint_memory_floor( "recompute_method": "uniform", "recompute_num_layers": 1, "distribute_saved_activations": False, - "sequence_parallel": tp > 1, + "sequence_parallel": sequence_parallel, "fp32_residual_connection": False, "cpu_offloading": False, "cuda_graph_impl": "none", @@ -4343,7 +4551,11 @@ def _checkpoint_memory_floor( or self._param_dtype_size != 2 or next(self.runtime.model[0].parameters()).dtype is not torch.bfloat16 or pp != 1 - or (tp > 1 and not self._sequence_parallel_floor_covered(layers, tp, cp)) + or ( + not (tp > 1 and self._sequence_parallel_floor_covered(layers, tp, cp)) + if sequence_parallel + else tp != 1 + ) or any( type(getattr(config, name, None)) is not type(value) or getattr(config, name) != value @@ -4358,7 +4570,48 @@ def _checkpoint_memory_floor( or getattr(decoder, "_forward_hooks", None) or getattr(decoder, "_forward_pre_hooks", None) ): + return None + return decoder + + def _checkpoint_memory_floor( + self, + group_rows: tuple[tuple[int, bool], ...], + slot_refs: tuple["LoRASlotRef | None", ...] | None = None, + gdn_segments: int = 0, + routed_rows: tuple[int, ...] | None = None, + layouts: tuple[_GroupLayout, ...] | None = None, + ) -> tuple[int, int]: + """Conservative saved-boundary charge and one recomputed layer's workspace. + + ``routed_rows`` are each group's balanced dispatched rows per rank + (``_plan_group_routed_rows``); by default, its local rows. With + ``layouts`` (``_plan_group_layouts``), price every rank on its own CP + layouts instead of the busiest rank's rows (``_layout_checkpoint_floor``). + Count actual local full/uniform/1 boundaries, including aliases, rather + than claiming measured distinct storage. Only this call's new groups + enter the term; already-live graphs remain in the availability baseline. + 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). + """ + gradient_rows = sum(rows for rows, grad in group_rows if grad) + _, tp, _, _ = self._topology_key() + decoder = ( + self._checkpoint_floor_decoder(sequence_parallel=tp > 1) + if group_rows + else None + ) + if decoder is None: return 0, 0 + layers = len(decoder.layers) refs = (None,) * len(group_rows) if slot_refs is None else slot_refs routed = (None,) * len(group_rows) if routed_rows is None else routed_rows if tp > 1: @@ -4437,13 +4690,18 @@ def _generic_checkpoint_floor( gradient_rows = sum(rows for rows, grad in group_rows if grad) retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 + dense, no_grad = self._dense_mlp_widths(refs) + dense = dense 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. + # and pre-MLP norm output, and its MoE stage its routing state; a + # covered dense layer keeps its MLP stage. mixer = ( self._recomputed_mixer_bytes_per_token() + ( 2 * self._hidden_size * 2 + self._moe_checkpoint_state_bytes_per_token() if moe + else 2 * self._hidden_size * 2 + dense + if dense else 0 ) if gradient_rows @@ -4453,12 +4711,12 @@ def _generic_checkpoint_floor( 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) + + (mixer * rows if grad else rows * max(no_grad, 4 * self._hidden_size * 2)) for (rows, grad), ref, dispatched in zip( group_rows, refs, routed, strict=True ) ) - if moe: + if moe or dense or no_grad: workspace += self._te_workspace_growth_bytes() return retained, workspace @@ -4492,7 +4750,14 @@ def _layout_checkpoint_floor( attention_inputs = len(layers) - gdn_inputs widths = self._recomputed_mixer_widths(stage_buffers=False) moe = self._checkpoint_moe_bytes_per_token() - beside = 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe else 0 + dense, _ = self._dense_mlp_widths(refs) + beside = ( + 2 * hidden + self._moe_checkpoint_state_bytes_per_token() + if moe + else 2 * hidden + dense + if dense + else 0 + ) retained_by_rank: list[int] = [] totals: list[int] = [] for rank in range(len(layouts[0].attention_rows)): @@ -4526,7 +4791,7 @@ def _layout_checkpoint_floor( totals.append(retained + workspace) retained = max(retained_by_rank) workspace = max(totals) - retained - if moe: + if moe or dense: workspace += self._te_workspace_growth_bytes() return retained, workspace @@ -4745,14 +5010,18 @@ def _checkpoint_input_gradient_bytes( 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. + other recompute work the floor does not price. A covered dense model + (every layer the traced gated MLP, ``_dense_mlp_widths``) also holds + one: Qwen3.8-27B CP2 traces show one H-wide input gradient at the peak. """ retained, _ = self._checkpoint_memory_floor(group_rows) if not retained: return 0 gradient_rows = sum(rows for rows, grad in group_rows if grad) refs = (None,) * len(group_rows) if slot_refs is None else slot_refs - if ( + # The same slots as the floor: if any group's slot falls back there, + # the per-boundary allowance must stay here too. + if self._dense_mlp_widths(refs)[0] or ( self._topology_key()[2] <= 2 and self._checkpoint_moe_bytes_per_token() and all( @@ -4764,6 +5033,40 @@ def _checkpoint_input_gradient_bytes( return gradient_rows * self._hidden_size * 2 return retained + def _dense_mlp_widths( + self, slot_refs: Sequence["LoRASlotRef | None"] | None = None + ) -> tuple[int, int]: + """The covered dense (gradient stage, no-grad transient) per row, or 0s. + + Only at CP2, where it was traced: at CP1 the attention allowance's + slack is smaller, and above CP2 a rank's remote attention stages may + keep more than the CP2 allowance; the per-boundary gradient allowance + still covers both. Named slots are rechecked: their adapters must stay + within the priced rank. Only where the checkpoint floor prices the + decoder, which also carries the TE workspace growth. + """ + stage = getattr(self, "_dense_recompute_bytes_per_token", 0) + no_grad = getattr(self, "_dense_no_grad_bytes_per_token", 0) + if ( + type(stage) is not int + or type(no_grad) is not int + or stage <= 0 + or no_grad <= 0 + or self._topology_key()[2] != 2 + or self._checkpoint_floor_decoder() is None + ): + return 0, 0 + for ref in slot_refs or (): + if ref is None or ref.name is None: + continue + slot = _dense_mlp_recompute_bytes_per_token( + self.runtime.model, ref, hidden_size=self._hidden_size + ) + if not all(slot): + return 0, 0 + stage, no_grad = max(stage, slot[0]), max(no_grad, slot[1]) + return stage, no_grad + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -4797,7 +5100,9 @@ def _subforward_cost( checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, hybridep_growth_bytes: int = 0, + lower_bound: bool = False, ) -> _SubforwardCost: + """``lower_bound`` prices optimistic rows from below (split pruning).""" required = self._estimate_required_memory_bytes_from_values( packed_tokens=packed_tokens, output_bytes=output_bytes, @@ -4812,6 +5117,7 @@ def _subforward_cost( checkpoint_floor=checkpoint_floor, retained_tokens=retained_tokens, include_checkpoint_input_gradient=False, + lower_bound=lower_bound, ) checkpoint_retained, checkpoint_workspace = self._checkpoint_memory_floor( group_rows, @@ -6990,6 +7296,8 @@ def _fill_planner_snapshot( "gdn_layers", "checkpointed_moe_layers", "moe_output_bytes_per_token", + "dense_recompute_bytes_per_token", + "dense_no_grad_bytes_per_token", ) } rank_fields["recompute_modules"] = sorted(self._recompute_modules) @@ -8532,6 +8840,7 @@ def _estimate_required_memory_bytes_from_values( checkpoint_floor: tuple[int, int] = (0, 0), retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, + lower_bound: bool = False, ) -> int: if packed_tokens <= 0: return output_bytes @@ -8543,6 +8852,29 @@ def _estimate_required_memory_bytes_from_values( * self._param_dtype_size * activation_factor ) + _dense_stage, no_grad = ( + self._dense_mlp_widths(slot_refs) + if not signature.grad_enabled + and signature.topology[2] == 2 + and len(group_rows) > 1 + else (0, 0) + ) + if no_grad: + # No-grad groups run one after another and keep only their outputs + # (charged below): price the largest group's own physical rows at + # the traced width. The per-packed-token floor is kept only as that + # group's share, which it matched for one group on Qwen3.8-27B. + # That share falls as another group's rows grow, so a lower bound + # on optimistic rows keeps only the largest group's own rows. + rows = [rows for rows, _ in group_rows] + static_compute = ( + max(rows) * no_grad + if lower_bound + else max( + max(rows) * no_grad, + -(-static_compute * max(rows) // max(1, sum(rows))), + ) + ) if signature.grad_enabled and self._recompute_granularity != "full": geometry = self._geometry hidden = self._hidden_size diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index 61b23fc8c..ec3b39d64 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -526,6 +526,10 @@ def report( "checkpointed_moe_layers recompute_modules moe_output_bytes_per_token " "moe_forward_stages".split() ) +# Recorded by newer ranks; reports from before it replay with 0 (no dense stage). +_OPTIONAL_RANK_FIELDS = frozenset( + {"dense_recompute_bytes_per_token", "dense_no_grad_bytes_per_token"} +) def _signature_values(values: dict[str, Any]) -> dict[str, Any]: @@ -587,13 +591,15 @@ def replay( if not state["estimates"]: raise ValueError("memory replay has no candidate estimates") values = state["rank"] - if set(values) != _RANK_FIELDS | {"geometry", "topology"}: + if set(values) - _OPTIONAL_RANK_FIELDS != _RANK_FIELDS | {"geometry", "topology"}: raise ValueError( "incomplete replay: immutable rank fields differ (including MoE stages)" ) rank = _impl.TrainerRank.__new__(_impl.TrainerRank) for name in _RANK_FIELDS - {"one_layer_recompute"}: setattr(rank, "_" + name, values[name]) + for name in _OPTIONAL_RANK_FIELDS: + setattr(rank, "_" + name, values.get(name, 0)) if type(values["one_layer_recompute"]) is not bool: raise ValueError("incomplete replay: recompute mode is not recorded") rank._recorded_one_layer_recompute = values["one_layer_recompute"] diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py new file mode 100644 index 000000000..e4e59c0c0 --- /dev/null +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -0,0 +1,673 @@ +"""Dense recompute and no-grad group pricing; CPU admission math, not a bound. + +Qwen3.8-27B CP2 allocator traces: the recomputed layer's peak holds its MLP +FC1 stage and one input gradient, and no-grad groups run one after another. +""" + +from dataclasses import replace +from types import MethodType, SimpleNamespace +from typing import Any + +import pytest +from test_trainer_rank_checkpoint_memory import price, rank, requests +from test_trainer_rank_moe_memory import _rank as _moe_rank +from test_trainer_rank_moe_memory import layer # noqa: F401 +import torch + +from art.trainer_rank import ForwardInput, _impl +from art.trainer_rank._impl import ( + Unset, + _dense_mlp_recompute_bytes_per_token, + _GroupLayout, + _MemorySignature, +) + +HIDDEN, LAYERS, FFN, RANK = 2048, 40, 5632, 8 +CP2 = (1, 1, 2, 1) +# 7F traced FC1 stage + one more FC1 triplet (6F) of early-run recompile +# residue, plus the adapters' rank-wide intermediates. +STAGE = (13 * FFN + 6 * RANK) * 2 +# Three 2F FC1 tensors, residual/norm/CP-gather rows, rank intermediates. +NO_GRAD = (6 * FFN + 6 * HIDDEN + 6 * RANK) * 2 + + +def _module(cls): + value = cls.__new__(cls) + torch.nn.Module.__init__(value) + return value + + +def _adapter(lora_module, inputs: int, outputs: int, rank: int = RANK): + lora = _module(lora_module.LoRA) + lora.A_T = torch.nn.Parameter(torch.empty(inputs, rank, dtype=torch.bfloat16)) + lora.B_T = torch.nn.Parameter(torch.empty(rank, outputs, dtype=torch.bfloat16)) + return lora + + +def _dense_layer(gdn: bool = False) -> Any: + """The traced gated MLP, from the real owner types.""" + pytest.importorskip("art.megatron.lora") + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TERowParallelLinear, + ) + from megatron.core.ssm.gated_delta_net import GatedDeltaNet + from megatron.core.transformer.attention import SelfAttention + from megatron.core.transformer.mlp import MLP + from megatron.core.transformer.transformer_layer import TransformerLayer + + from art.megatron import lora as lora_module + + mlp = _module(MLP) + mlp.config = SimpleNamespace( + hidden_size=HIDDEN, + ffn_hidden_size=FFN, + gated_linear_unit=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + sequence_parallel=False, + bias_activation_fusion=True, # Bridge's Qwen3.5 providers fuse SwiGLU. + use_te_activation_func=False, + cpu_offloading=False, + cuda_graph_impl="none", + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + fp8=None, + fp4=None, + activation_func=torch.nn.functional.silu, + activation_func_clamp_value=None, + glu_linear_offset=0.0, + ) + mlp.activation_func = torch.nn.functional.silu + fc1 = _module(lora_module.SharedExpertsLinearFC1LoRA) + fc1.linear_fc1 = _module(TELayerNormColumnParallelLinear) + fc1.gate_lora = _adapter(lora_module, HIDDEN, FFN) + fc1.up_lora = _adapter(lora_module, HIDDEN, FFN) + fc1.non_gated = False + fc1.out_features = 2 * FFN + fc2 = _module(lora_module.SharedExpertsLinearFC2LoRA) + row = _module(lora_module.SelfAttentionLinearProjLoRA) + row.lora = _adapter(lora_module, FFN, HIDDEN) + row.linear_proj = _module(TERowParallelLinear) + fc2.row_parallel_lora = row + mlp.linear_fc1 = fc1 + mlp.linear_fc2 = fc2 + layer = _module(TransformerLayer) + layer.self_attention = _module(GatedDeltaNet if gdn else SelfAttention) + layer.mlp = mlp + return layer + + +def _wrap_like_art(layer: Any) -> None: + """ART's GDN island and prefix-tree wrappers, as the traced run had them.""" + from art.megatron.gdn.operator import ( + _gdn_island_layer_forward, + _prefix_tree_forward, + ) + + # Training compile then wraps the delegate (training/compile.py). + layer._art_gdn_island_physical_forward = torch.compile(layer.forward) + layer.forward = MethodType(_gdn_island_layer_forward, layer) + mixer = layer.self_attention + if type(mixer).__name__ == "GatedDeltaNet": + mixer._art_physical_forward = mixer.forward + mixer.forward = MethodType(_prefix_tree_forward, mixer) + + +def _dense_model(layers: list[Any]) -> Any: + from megatron.core.transformer.transformer_block import TransformerBlock + + block = _module(TransformerBlock) + block.layers = torch.nn.ModuleList(layers) + model: Any = torch.nn.Module() + model.decoder = block + model._preprocess = lambda: None # Marks a GPT model for _language_model. + return model + + +def _dense_rank(stage: int = STAGE, no_grad: int = NO_GRAD): + r = rank() + r._moe_output_bytes_per_token = r._moe_checkpoint_grad_bytes_per_token = 0 + r._moe_recompute_covered = False + r._dense_recompute_bytes_per_token = stage + r._dense_no_grad_bytes_per_token = no_grad + return r + + +def _at_cp2(r): + # After cheap estimation (which declines under CP): price as a CP2 rank. + r._topology_key = lambda: CP2 + return r + + +def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): + # A GDN/attention hybrid with ART's wrappers, as traced. + layers = [_dense_layer(gdn=index != 2) for index in range(3)] + for layer in layers: + _wrap_like_art(layer) + model = _dense_model(layers) + assert _dense_mlp_recompute_bytes_per_token([model]) == (STAGE, NO_GRAD) + assert _dense_mlp_recompute_bytes_per_token([model], hidden_size=HIDDEN + 1) == ( + 0, + 0, + ) + # Qwen3.8-27B at rank 8: F 17,408, H 5,120. + assert (13 * 17408 + 48) * 2 == 452_704 + + +@pytest.mark.parametrize( + "change", + [ + "fc1_hook", + "row_hook", + "lora_hook", + "base_hook", + "layer_hook", + "layer_forward", + "layer_delegate", + "compiled_custom_delegate", + "layer_type", + "mixer_delegate", + "mixer_class_forward", + "mixer_hook", + "mixer_forward", + "no_mixer", + "decoder_hook", + "decoder_subclass", + "non_gated", + "activation", + "config_activation", + "clamp", + "linear_offset", + "unfused_activation", + "te_activation", + "fp8", + "fp4", + "dtype", + "tensor_parallel", + "pipeline_parallel", + "sequence_parallel", + "cuda_graphs", + "bias", + "out_features", + "other_layer", + "unwrapped_fc1", + "unfused_norm", + "rank", + "mixer_adapter_rank", + "mixer_child_hook", + "mixer_child_forward", + "norm_delegate", + "slot_selector", + "active_selector", + "chunks", + ], +) +def test_anything_but_the_traced_execution_keeps_the_allowance(change): + from megatron.core.extensions.transformer_engine import TEColumnParallelLinear + from megatron.core.transformer.attention import SelfAttention + from megatron.core.transformer.transformer_block import TransformerBlock + from megatron.core.transformer.transformer_layer import TransformerLayer + + from art.megatron import lora as lora_module + from art.megatron.gdn.operator import _empty_safe_norm_forward + + layers = [_dense_layer(gdn=index != 1) for index in range(3)] + for wrapped in layers: + _wrap_like_art(wrapped) + model = _dense_model(layers) + layer = layers[1] + gdn_layer = layers[0] + mlp = layer.mlp + config = mlp.config + + def hook(module): + return lambda: module.register_forward_hook(lambda *args: None) + + class Block(TransformerBlock): + pass + + class Layer(TransformerLayer): + pass + + class Attention(SelfAttention): + def forward(self, *args, **kwargs): # A class-level override. + return super().forward(*args, **kwargs) + + custom = MethodType(lambda self, *a, **k: None, layer) + + def mixer_child(edit): + # A child the mixer runs, such as its core attention or a norm. + def apply(): + child = torch.nn.LayerNorm(4) + edit(child) + layer.self_attention.core_attention = child + + return apply + + def foreign_norm_delegate(norm): + norm._art_empty_safe_norm_physical_forward = MethodType(lambda self, x: x, norm) + norm.forward = MethodType(_empty_safe_norm_forward, norm) + + edits = { + "fc1_hook": hook(mlp.linear_fc1), + "row_hook": hook(mlp.linear_fc2.row_parallel_lora), + "lora_hook": hook(mlp.linear_fc1.gate_lora), + "base_hook": hook(mlp.linear_fc1.linear_fc1), + "layer_hook": lambda: layer.register_forward_pre_hook(lambda *args: None), + "layer_forward": lambda: setattr( + layer, "forward", MethodType(lambda self, *a: None, layer) + ), + "layer_delegate": lambda: setattr( + layer, "_art_gdn_island_physical_forward", custom + ), + "compiled_custom_delegate": lambda: setattr( + layer, "_art_gdn_island_physical_forward", torch.compile(custom) + ), + "layer_type": lambda: setattr(layer, "__class__", Layer), + "mixer_delegate": lambda: setattr( + gdn_layer.self_attention, + "_art_physical_forward", + MethodType(lambda self, *a: None, gdn_layer.self_attention), + ), + "mixer_class_forward": lambda: setattr( + layer.self_attention, "__class__", Attention + ), + "mixer_hook": hook(layer.self_attention), + "mixer_forward": lambda: setattr( + layer.self_attention, "forward", lambda *a: None + ), + "no_mixer": lambda: delattr(layer, "self_attention"), + "decoder_hook": hook(model.decoder), + "decoder_subclass": lambda: setattr(model.decoder, "__class__", Block), + "non_gated": lambda: setattr(mlp.linear_fc1, "non_gated", True), + "activation": lambda: setattr(mlp, "activation_func", torch.nn.functional.gelu), + "config_activation": lambda: setattr( + config, "activation_func", torch.nn.functional.gelu + ), + "clamp": lambda: setattr(config, "activation_func_clamp_value", 7.0), + "linear_offset": lambda: setattr(config, "glu_linear_offset", 1.0), + "unfused_activation": lambda: setattr(config, "bias_activation_fusion", False), + "te_activation": lambda: setattr(config, "use_te_activation_func", True), + "fp8": lambda: setattr(config, "fp8", "hybrid"), + "fp4": lambda: setattr(config, "fp4", "nvfp4"), + "dtype": lambda: setattr(config, "params_dtype", torch.float16), + "tensor_parallel": lambda: setattr(config, "tensor_model_parallel_size", 2), + "pipeline_parallel": lambda: setattr(config, "pipeline_model_parallel_size", 2), + "sequence_parallel": lambda: setattr(config, "sequence_parallel", True), + "cuda_graphs": lambda: setattr(config, "cuda_graph_impl", "local"), + "bias": lambda: setattr(config, "add_bias_linear", True), + "out_features": lambda: setattr(mlp.linear_fc1, "out_features", FFN), + "other_layer": lambda: setattr(layer, "mlp", torch.nn.Linear(1, 1)), + "unwrapped_fc1": lambda: setattr(mlp, "linear_fc1", mlp.linear_fc1.linear_fc1), + "unfused_norm": lambda: setattr( + mlp.linear_fc1, "linear_fc1", _module(TEColumnParallelLinear) + ), + "rank": lambda: setattr( + mlp.linear_fc1, "up_lora", _adapter(lora_module, HIDDEN, FFN, 512) + ), + "mixer_adapter_rank": lambda: setattr( + layer.self_attention, "qkv_lora", _adapter(lora_module, HIDDEN, HIDDEN, 512) + ), + "mixer_child_hook": mixer_child( + lambda child: child.register_forward_hook(lambda *args: None) + ), + "mixer_child_forward": mixer_child( + lambda child: setattr(child, "forward", MethodType(lambda s, x: x, child)) + ), + "norm_delegate": mixer_child(foreign_norm_delegate), + # Execution selects tensors through the instance; the gate must too. + "slot_selector": lambda: setattr( + mlp.linear_fc1.gate_lora, "_slot", lambda ref: None + ), + "active_selector": lambda: setattr( + mlp.linear_fc1.gate_lora, "active_lora_tensors", lambda: None + ), + "chunks": lambda: None, + } + edits[change]() + models = [model] * (2 if change == "chunks" else 1) + assert _dense_mlp_recompute_bytes_per_token(models) == (0, 0) + + +def test_moe_models_never_price_the_dense_stage(monkeypatch, layer): + calls = [] + monkeypatch.setattr( + _impl, + "_dense_mlp_recompute_bytes_per_token", + lambda *a, **k: calls.append(1) or (STAGE, NO_GRAD), + ) + moe = _moe_rank(layer) + assert moe._moe_layers and not calls + assert ( + moe._dense_recompute_bytes_per_token, + moe._dense_no_grad_bytes_per_token, + ) == (0, 0) + + +def test_an_active_slot_beyond_the_priced_rank_keeps_the_allowance(monkeypatch): + r = _at_cp2(_dense_rank()) + policy = r._slot_ref("policy") + assert r._dense_mlp_widths((None,)) == (STAGE, NO_GRAD) + monkeypatch.setattr( + _impl, "_dense_mlp_recompute_bytes_per_token", lambda *a, **k: (0, 0) + ) + assert r._dense_mlp_widths((policy,)) == (0, 0) + wider = (STAGE + 96, NO_GRAD + 96) + monkeypatch.setattr( + _impl, "_dense_mlp_recompute_bytes_per_token", lambda *a, **k: wider + ) + assert r._dense_mlp_widths((policy,)) == wider + + +@pytest.mark.parametrize("rows", [67, 1024]) +def test_covered_dense_recompute_charges_one_gradient_and_its_stage(rows): + r = _dense_rank() + # A short no-grad reference keeps its transient below the gradient stage. + values = r._estimate_flat_forward(requests(rows, 16)) + _at_cp2(r) + cost = price(r, values) + assert cost.checkpoint_input_gradient == rows * HIDDEN * 2 + # The recomputed mixer keeps its activations beside the residual, the + # norm output and the MLP stage (with its recompile residue); with no + # learned profile the floor alone carries it. + assert r._memory_profiles.get(values[2]) is None + per_row = r._recomputed_mixer_bytes_per_token() + 2 * HIDDEN * 2 + STAGE + assert cost.checkpoint_workspace >= rows * per_row + r._te_workspace_growth_bytes() + assert cost.required >= int( + ( + cost.checkpoint_retained + + cost.checkpoint_workspace + + cost.checkpoint_input_gradient + ) + * 1.1 + ) + # Without the traced stage, dense keeps one gradient per boundary. + plain = price(_at_cp2(_dense_rank(0, 0)), values) + assert plain.checkpoint_input_gradient == rows * LAYERS * HIDDEN * 2 + + +def test_a_larger_no_grad_group_keeps_its_transient_beside_gradient_boundaries(): + r = _dense_rank() + values = r._estimate_flat_forward(requests(1000, 3000)) + cost = price(_at_cp2(r), values) + boundaries = 1000 * LAYERS * HIDDEN * 2 + # The later no-grad group's FC1 transient runs beside the retained graph. + assert cost.checkpoint_workspace >= 3000 * NO_GRAD + assert cost.required >= int((boundaries + 3000 * NO_GRAD + 1000 * HIDDEN * 2) * 1.1) + + +@pytest.mark.parametrize("topology", [(1, 1, 1, 1), (1, 1, 4, 1)]) +def test_dense_widths_apply_only_at_cp2(topology): + r = _dense_rank() + values = r._estimate_flat_forward(requests(67, 16)) + r._topology_key = lambda: topology + assert r._dense_mlp_widths() == (0, 0) + assert price(r, values).checkpoint_input_gradient == 67 * LAYERS * HIDDEN * 2 + + +@pytest.mark.parametrize("sequence_parallel", [False, True]) +@pytest.mark.parametrize("tp", [2, 4]) +def test_dense_widths_stay_at_tp1(monkeypatch, tp, sequence_parallel): + """Even where a TP x SP floor prices CP2, the CP2 dense trace stays TP1.""" + r = _dense_rank() + values = r._estimate_flat_forward(requests(67, 16)) + r._topology_key = lambda: (1, tp, 2, 1) + decoder = _impl._language_model(r.runtime.model[0]).decoder + decoder.config.sequence_parallel = sequence_parallel + monkeypatch.setattr(r, "_sequence_parallel_floor_covered", lambda *_: True) + assert r._checkpoint_floor_decoder() is None + assert r._dense_mlp_widths() == (0, 0) + if sequence_parallel: + assert r._checkpoint_floor_decoder(sequence_parallel=True) is decoder + # The floor prices it, but with one gradient per sharded boundary. + rows = -(-67 // tp) + assert price(r, values).checkpoint_input_gradient == rows * LAYERS * HIDDEN * 2 + + +@pytest.mark.parametrize("gdn", [False, True]) +def test_layout_floor_prices_the_dense_stage_per_layer_type(gdn): + r = _at_cp2(_dense_rank()) + plain = _at_cp2(_dense_rank(0, 0)) + if gdn: + for x in (r, plain): + x._gdn_layers = 3 + x._geometry = replace( + x._geometry, + gdn_key_heads=4, + gdn_key_head_dim=64, + gdn_value_heads=8, + gdn_value_head_dim=64, + ) + layers = [ + SimpleNamespace( + _art_gdn_island_boundary=SimpleNamespace( + input_layout="gdn" if gdn and index else "attention" + ) + ) + for index in range(4) + ] + layout = _GroupLayout( + attention_rows=(100, 80), + gdn_rows=(120, 60) if gdn else None, + attention_retained=(0, 0), + ) + dense = r._layout_checkpoint_floor(layers, (None,), (None,), (layout,)) + base = plain._layout_checkpoint_floor(layers, (None,), (None,), (layout,)) + assert dense[0] == base[0] + # Every layer type's recomputed stage gains the residual, norm and MLP + # stage on its own rows: the busiest rank's largest stage grows by it. + widths = r._recomputed_mixer_widths(stage_buffers=False) + rows = {"attention": 100, "gdn": 120} if gdn else {"attention": 100} + grown = max(rows[k] * (widths[k] + 2 * HIDDEN * 2 + STAGE) for k in rows) + plain_stage = max(rows[k] * widths[k] for k in rows) + assert ( + sum(dense) - sum(base) == grown - plain_stage + r._te_workspace_growth_bytes() + ) + + +def _no_grad_required(r, group_rows, *, topology=CP2, packed=40_000): + signature = _MemorySignature( + topology, (1, None), len(group_rows), (), False, (False,) * len(group_rows) + ) + return r._estimate_required_memory_bytes_from_values( + packed_tokens=packed, + output_bytes=0, + signature=signature, + logical_tokens=packed, + group_rows=tuple((rows, False) for rows in group_rows), + ) + + +def test_no_grad_groups_price_the_largest_groups_own_rows(): + dense, plain = _at_cp2(_dense_rank()), _at_cp2(_dense_rank(0, 0)) + # Today's floor: H bytes per packed token times the layer-count factor. + per_token = HIDDEN * 2 * min(16, LAYERS // 4 + 4) + te = dense._te_workspace_growth_bytes() + # One group: today's per-packed-token floor, unchanged. + assert _no_grad_required(dense, (12_000,)) == _no_grad_required(plain, (12_000,)) + for groups, packed in (((12_000, 8_000), 40_000), ((58_240, 29_120), 119_119)): + largest = max(groups) + # The largest group's own physical rows at the traced width (with + # TE's workspace growth), and at least its share of today's + # per-packed-token floor. + expected = max( + largest * NO_GRAD + te, -(-packed * per_token * largest // sum(groups)) + ) + assert _no_grad_required(dense, groups, packed=packed) == int(expected * 1.1) + assert _no_grad_required(dense, groups, packed=packed) < _no_grad_required( + plain, groups, packed=packed + ) + # A narrow structural width never drops below the rows' traced need. + wide = _at_cp2(_dense_rank(STAGE, 10**6)) + assert _no_grad_required(wide, (12_000, 8_000)) == int((12_000 * 10**6 + te) * 1.1) + + +@pytest.mark.parametrize("case", ["cp1", "cp4", "unsupported"]) +def test_no_grad_group_floor_needs_the_traced_shape(case): + topology = {"cp1": (1, 1, 1, 1), "cp4": (1, 1, 4, 1)}.get(case, CP2) + dense = _dense_rank(*((0, 0) if case == "unsupported" else (STAGE, NO_GRAD))) + plain = _dense_rank(0, 0) + for r in (dense, plain): + r._topology_key = lambda: topology + assert _no_grad_required( + dense, (12_000, 8_000), topology=topology + ) == _no_grad_required(plain, (12_000, 8_000), topology=topology) + + +def test_one_unsupported_slot_keeps_both_allowances(monkeypatch): + """A no-grad reference slot outside the priced rank must not leave the + gradient group's boundaries on one gradient without the dense stage.""" + r = _at_cp2(_dense_rank()) + policy, reference = r._slot_ref("policy"), r._slot_ref("reference") + monkeypatch.setattr( + _impl, + "_dense_mlp_recompute_bytes_per_token", + lambda model, ref, **k: (0, 0) if ref == reference else (STAGE, NO_GRAD), + ) + groups = ((1000, True), (3000, False)) + both = (policy, reference) + assert ( + r._checkpoint_input_gradient_bytes(groups, both) == 1000 * LAYERS * HIDDEN * 2 + ) + retained, workspace = r._checkpoint_memory_floor(groups, both) + assert workspace < 3000 * NO_GRAD # the dense widths fell back together + supported = (policy, policy) + assert r._checkpoint_input_gradient_bytes(groups, supported) == 1000 * HIDDEN * 2 + + +def test_no_grad_only_waves_charge_te_workspace_growth(): + r = _at_cp2(_dense_rank()) + _, dense = r._checkpoint_memory_floor(((500, False),)) + _, plain = _at_cp2(_dense_rank(0, 0))._checkpoint_memory_floor(((500, False),)) + assert dense == 500 * NO_GRAD + r._te_workspace_growth_bytes() + assert plain == 500 * 4 * HIDDEN * 2 + + +def test_the_traced_qwen_attention_mixer_is_accepted(): + bridge = pytest.importorskip( + "megatron.bridge.models.qwen_vl.modelling_qwen3_vl.attention" + ) + layers = [_dense_layer(gdn=index != 1) for index in range(3)] + qwen = _module(bridge.Qwen3VLSelfAttention) + layers[1].self_attention = qwen # Overrides forward, as traced. + for layer in layers: + _wrap_like_art(layer) + assert _dense_mlp_recompute_bytes_per_token([_dense_model(layers)]) == ( + STAGE, + NO_GRAD, + ) + + +def test_every_adapter_the_layer_runs_is_priced_beside_arts_norm_wrapper(): + from art.megatron import lora as lora_module + from art.megatron.gdn.operator import _empty_safe_norm_forward + + layers = [_dense_layer(gdn=index != 2) for index in range(3)] + for layer in layers: + _wrap_like_art(layer) + # A mixer adapter keeps its rank-wide products beside the MLP's. + layers[1].self_attention.qkv_lora = _adapter(lora_module, HIDDEN, HIDDEN, 32) + # ART's empty-safe norm wrapper still calls the norm's own forward. + norm = torch.nn.LayerNorm(HIDDEN) + norm._art_empty_safe_norm_physical_forward = norm.forward + norm.forward = MethodType(_empty_safe_norm_forward, norm) + layers[0].self_attention.q_layernorm = norm + ranks = 2 * (3 * RANK + 32) + assert _dense_mlp_recompute_bytes_per_token([_dense_model(layers)]) == ( + (13 * FFN + ranks) * 2, + (6 * FFN + 6 * HIDDEN + ranks) * 2, + ) + + +def test_named_slots_are_read_through_the_lookup_execution_uses(): + from art.megatron import lora as lora_module + + layers = [_dense_layer(gdn=index == 0) for index in range(2)] + for layer in layers: + _wrap_like_art(layer) + model = _dense_model(layers) + policy = rank()._slot_ref("policy") + adapters = [ + m for layer in layers for m in layer.modules() if type(m) is lora_module.LoRA + ] + + def load(adapter, width): + slot = _module(lora_module.LoRASlot) + slot.A_T = torch.nn.Parameter( + torch.empty(adapter.A_T.shape[0], width, dtype=torch.bfloat16) + ) + slot.B_T = torch.nn.Parameter( + torch.empty(width, adapter.B_T.shape[1], dtype=torch.bfloat16) + ) + adapter._slot_keys = {policy: "slot_0"} + adapter._slot_modules = torch.nn.ModuleDict({"slot_0": slot}) + + for adapter in adapters: + load(adapter, 16) + widths = (13 * FFN + 2 * 3 * 16) * 2, (6 * FFN + 6 * HIDDEN + 2 * 3 * 16) * 2 + assert _dense_mlp_recompute_bytes_per_token([model], policy) == widths + # A slot without an adapter on one module runs the base output there. + for layer in layers: + layer.mlp.linear_fc1.up_lora._slot_keys = {} + assert _dense_mlp_recompute_bytes_per_token([model], policy) == ( + (13 * FFN + 2 * 2 * 16) * 2, + (6 * FFN + 6 * HIDDEN + 2 * 2 * 16) * 2, + ) + load(adapters[0], 300) # Loaded wider than the priced rank. + assert _dense_mlp_recompute_bytes_per_token([model], policy) == (0, 0) + + +@pytest.mark.parametrize("case", ["selective", "eval"]) +def test_dense_widths_need_the_checkpoint_floors_decoder(case): + """The no-grad discount must not outlive the floor that adds TE growth.""" + r = _at_cp2(_dense_rank()) + assert r._dense_mlp_widths() == (STAGE, NO_GRAD) + decoder = _impl._language_model(r.runtime.model[0]).decoder + if case == "selective": + decoder.config.recompute_granularity = "selective" + else: + decoder.train(False) + assert r._checkpoint_memory_floor(((12_000, False), (8_000, False))) == (0, 0) + assert r._dense_mlp_widths() == (0, 0) + discounted = _no_grad_required(r, (12_000, 8_000)) + r._dense_recompute_bytes_per_token = r._dense_no_grad_bytes_per_token = 0 + assert _no_grad_required(r, (12_000, 8_000)) == discounted + + +def test_the_split_lower_bound_never_exceeds_the_exact_no_grad_price(monkeypatch): + r = _at_cp2(_dense_rank()) + signature = _MemorySignature(CP2, (1, None), 2, (), False, (False, False)) + + def required(rows, lower_bound): + return r._subforward_cost( + packed_tokens=20_480, + output_bytes=0, + signature=signature, + logical_tokens=20_480, + group_rows=tuple((n, False) for n in rows), + lower_bound=lower_bound, + ).required + + # Even shares bound the busiest rank's rows from below, but the largest + # group's share of today's floor falls as the other group's rows grow. + assert required((8192, 2048), False) > required((8192, 4096), False) + assert required((8192, 2048), True) <= required((8192, 4096), False) + # The split planner prices its optimistic rows in that mode. + modes = [] + exact = r._subforward_cost + monkeypatch.setattr( + r, + "_subforward_cost", + lambda **kwargs: modes.append(kwargs.get("lower_bound")) or exact(**kwargs), + ) + chunk = [ + ForwardInput( + input_tokens=torch.arange(64), target_tokens=torch.arange(64), no_grad=True + ) + for _ in range(2) + ] + r._split_chunk_lower_cost( + chunk, tuple(q.input_tokens for q in chunk), checkpoint=Unset + ) + assert modes == [True] diff --git a/tests/unit/test_trainer_rank_planner_reports.py b/tests/unit/test_trainer_rank_planner_reports.py index 1adce5472..252c0b34f 100644 --- a/tests/unit/test_trainer_rank_planner_reports.py +++ b/tests/unit/test_trainer_rank_planner_reports.py @@ -235,7 +235,8 @@ def test_spool_symlink_refuses(tmp_path): assert not list(target.iterdir()) -def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): +@pytest.mark.parametrize("dense_field", [False, True]) +def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path, dense_field): from art.trainer_rank._prefix_tree_planner import ( build_canonical_prefix_tree, plan_prefix_tree_layout, @@ -259,6 +260,8 @@ def test_replay_reruns_real_memory_estimator_and_prefix_layout(tmp_path): "recompute_modules": [], "moe_output_bytes_per_token": 0, "moe_forward_stages": [], + # Reports from before the dense stage field replay without it. + **({"dense_recompute_bytes_per_token": 0} if dense_field else {}), "geometry": { "hidden_size": 8, "ffn_hidden_size": 32,