From 29024e3d6e2f552a3fa44ec6cd75468619cf9dc5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:20:42 +0000 Subject: [PATCH 01/16] Price dense recompute by its traced stage and one input gradient For dense models TrainerRank charged one input gradient per saved boundary, which on dense Qwen3.8-27B at CP2 over-prices gradient waves by about 40%. Allocator traces show the recomputed layer's peak holds one input gradient and an MLP FC1 stage of 7F per row: the base output, the LoRA output and their sum, plus one F-wide tensor. When every decoder layer is the exact supported gated MLP, up to CP2, price that stage beside the recomputed mixer in both checkpoint floors and charge one gradient. Above CP2, or for any other structure, the per-boundary allowance stays. No-grad groups run one after another, so a CP2 multi-group no-grad wave's static floor now counts only the largest group's share of packed tokens. Single-group pricing is unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 4 +- src/art/trainer_rank/_impl.py | 151 +++++++++++- src/art/trainer_rank/_planner_misses.py | 6 +- tests/unit/test_trainer_rank_dense_memory.py | 220 ++++++++++++++++++ .../unit/test_trainer_rank_planner_reports.py | 5 +- 5 files changed, 375 insertions(+), 11 deletions(-) create mode 100644 tests/unit/test_trainer_rank_dense_memory.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index b66ac2c0e..8fa7b6363 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -241,6 +241,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 \ @@ -287,4 +288,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 6d3d61b98..6f205935c 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1630,6 +1630,92 @@ def _hybridep_buffer_bytes(capacity: int, ranks: int, hidden: int, experts: int) return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) +def _dense_mlp_recompute_bytes_per_token(model: Sequence[torch.nn.Module]) -> int: + """The dense MLP stage live at a recomputed layer's backward peak, per row. + + Qwen3.8-27B allocator traces at CP2 (dense, gated SwiGLU, LoRA on FC1 and + FC2): the peak sits in the recomputed layer's FC1 stage, holding the FC1 + base output, the LoRA gate/up output and their sum (2F each) plus one + F-wide tensor: 7F elements per row (238 KB measured, 244 KB priced at + F = 17,408). Every decoder layer must be this exact supported MLP; + otherwise 0 keeps the per-boundary gradient allowance. + """ + if len(model) != 1: + return 0 + try: + decoder = _language_model(model[0]).decoder + except (AttributeError, RuntimeError): + return 0 + layers = getattr(decoder, "layers", None) + if not layers or not all(hasattr(layer, "mlp") for layer in layers): + return 0 + try: + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + TERowParallelLinear, + ) + from megatron.core.transformer.mlp import MLP + from megatron.core.transformer.transformer_block import TransformerBlock + + from art.megatron.lora import ( + LoRA, + SelfAttentionLinearProjLoRA, + SharedExpertsLinearFC1LoRA, + SharedExpertsLinearFC2LoRA, + ) + except ImportError: + # Without the traced owner types nothing can match; keep the allowance. + return 0 + if type(decoder) is not TransformerBlock: + return 0 + stage = 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) + base1 = getattr(fc1, "linear_fc1", None) + sites = ( + (mlp, MLP), + (fc1, SharedExpertsLinearFC1LoRA), + (fc2, SharedExpertsLinearFC2LoRA), + (row, SelfAttentionLinearProjLoRA), + (getattr(row, "lora", None), LoRA), + (getattr(row, "linear_proj", None), TERowParallelLinear), + (getattr(fc1, "gate_lora", None), LoRA), + (getattr(fc1, "up_lora", None), LoRA), + ) + ffn = getattr(config, "ffn_hidden_size", None) + if ( + type(base1) not in (TEColumnParallelLinear, TELayerNormColumnParallelLinear) + or any(type(site) is not cls for site, cls in sites) + or any( + "forward" in vars(site) + or cast(Any, site)._forward_hooks + or cast(Any, site)._forward_pre_hooks + for site in (base1, *(site for site, _ in sites)) + ) + or type(ffn) is not int + or ffn <= 0 + or getattr(fc1, "non_gated", None) is not False + or getattr(fc1, "out_features", None) != 2 * ffn + or getattr(config, "gated_linear_unit", None) is not True + or getattr(config, "params_dtype", None) is not torch.bfloat16 + or getattr(config, "add_bias_linear", None) is not False + or getattr(config, "sequence_parallel", None) is not False + or getattr(config, "fp8", None) + or getattr(config, "fp4", None) + or getattr(config, "cuda_graph_impl", "none") != "none" + or getattr(config, "tensor_model_parallel_size", None) != 1 + or getattr(config, "pipeline_model_parallel_size", None) != 1 + or getattr(mlp, "activation_func", None) is not torch.nn.functional.silu + ): + return 0 + stage = max(stage, 7 * ffn * 2) + return stage + + def _moe_output_bytes_per_token( model: Sequence[torch.nn.Module], shape: ParallelShape, @@ -2089,6 +2175,13 @@ 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 = ( + 0 + if self._moe_layers + else _dense_mlp_recompute_bytes_per_token(runtime.model) + ) self._ep_group_is_cp_group = _ep_group_is_cp_group(self._parallel_shape) selection = select_scoring( device_capability=capability, @@ -4276,13 +4369,17 @@ def _checkpoint_memory_floor( return self._layout_checkpoint_floor(decoder.layers, refs, routed, layouts) retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 + dense = self._dense_recompute_stage_bytes() 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 @@ -4331,7 +4428,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_recompute_stage_bytes() + 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)): @@ -4583,21 +4687,38 @@ def _checkpoint_input_gradient_bytes( every layer's recompute, FC1 included (Qwen3.6-35B-A3B traces at CP1, CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere keep one gradient per boundary: that allowance also covers dense MLP and other recompute - work the floor does not price. + work the floor does not price. A covered dense model (every layer the + traced gated MLP, ``_dense_recompute_stage_bytes``) 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 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 + if self._dense_recompute_stage_bytes() or ( + 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 + ) ): return gradient_rows * self._hidden_size * 2 return retained + def _dense_recompute_stage_bytes(self) -> int: + """The covered dense MLP stage per recomputed row, where it was traced. + + Only up to CP2: at higher CP a rank's remote attention stages may keep + more than the CP2 allowance, which the per-boundary gradient allowance + still covers there. + """ + stage = getattr(self, "_dense_recompute_bytes_per_token", 0) + if type(stage) is not int or stage <= 0 or self._topology_key()[2] > 2: + return 0 + return stage + def _plan_cost(self, plan: _FlatForwardPlan) -> _SubforwardCost: return self._subforward_cost( packed_tokens=plan.packed_tokens, @@ -6787,6 +6908,7 @@ def _fill_planner_snapshot( "gdn_layers", "checkpointed_moe_layers", "moe_output_bytes_per_token", + "dense_recompute_bytes_per_token", ) } rank_fields["recompute_modules"] = sorted(self._recompute_modules) @@ -8334,8 +8456,21 @@ def _estimate_required_memory_bytes_from_values( return output_bytes profiled = self._memory_profiles.get(signature) activation_factor = max(4, min(16, self._num_layers // 4 + 4)) + floor_tokens = packed_tokens + if ( + not signature.grad_enabled + and signature.topology[2] == 2 + and len(group_rows) > 1 + and self._dense_recompute_stage_bytes() + ): + # No-grad groups run one after another and keep nothing but their + # outputs (charged below), so only the largest group's transient + # is live. Qwen3.8-27B CP2 traces: 263 KB per busiest-rank row of + # one group, which this floor matches for a single group. + rows = [rows for rows, _ in group_rows] + floor_tokens = -(-packed_tokens * max(rows) // max(1, sum(rows))) static_compute = ( - packed_tokens + floor_tokens * self._hidden_size * self._param_dtype_size * activation_factor diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index 61b23fc8c..566cdc126 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -526,6 +526,8 @@ 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"}) def _signature_values(values: dict[str, Any]) -> dict[str, Any]: @@ -587,13 +589,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..65bb04aa3 --- /dev/null +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -0,0 +1,220 @@ +"""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 (7F per row) and one input gradient, and no-grad groups run one +after another. +""" + +from types import SimpleNamespace +from typing import Any + +import pytest +from test_trainer_rank_checkpoint_memory import price, rank, requests +import torch + +from art.trainer_rank._impl import ( + _dense_mlp_recompute_bytes_per_token, + _GroupLayout, + _MemorySignature, +) + +HIDDEN, LAYERS, FFN = 2048, 40, 5632 +STAGE = 7 * FFN * 2 + + +def _module(cls): + value = cls.__new__(cls) + torch.nn.Module.__init__(value) + return value + + +def _dense_layer() -> Any: + """The supported gated MLP, from the real owner types.""" + pytest.importorskip("art.megatron.lora") + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TERowParallelLinear, + ) + from megatron.core.transformer.mlp import MLP + + from art.megatron import lora as lora_module + + mlp = _module(MLP) + mlp.config = SimpleNamespace( + ffn_hidden_size=FFN, + gated_linear_unit=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + sequence_parallel=False, + fp8=None, + fp4=None, + cuda_graph_impl="none", + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + mlp.activation_func = torch.nn.functional.silu + fc1 = _module(lora_module.SharedExpertsLinearFC1LoRA) + fc1.linear_fc1 = _module(TELayerNormColumnParallelLinear) + fc1.gate_lora = _module(lora_module.LoRA) + fc1.up_lora = _module(lora_module.LoRA) + fc1.non_gated = False + fc1.out_features = 2 * FFN + fc2 = _module(lora_module.SharedExpertsLinearFC2LoRA) + row = _module(lora_module.SelfAttentionLinearProjLoRA) + row.lora = _module(lora_module.LoRA) + row.linear_proj = _module(TERowParallelLinear) + fc2.row_parallel_lora = row + mlp.linear_fc1 = fc1 + mlp.linear_fc2 = fc2 + layer = torch.nn.Module() + layer.mlp = mlp + return layer + + +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): + 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 + return r + + +def test_supported_dense_mlp_prices_its_traced_stage(): + layers = [_dense_layer() for _ in range(3)] + assert _dense_mlp_recompute_bytes_per_token([_dense_model(layers)]) == STAGE + + +@pytest.mark.parametrize( + "change", + [ + "hook", + "non_gated", + "activation", + "fp8", + "tensor_parallel", + "bias", + "out_features", + "other_layer", + "unwrapped_fc1", + "chunks", + ], +) +def test_any_unsupported_layer_keeps_the_boundary_allowance(change): + layers = [_dense_layer() for _ in range(3)] + mlp = layers[1].mlp + if change == "hook": + mlp.linear_fc1.register_forward_hook(lambda *args: None) + elif change == "non_gated": + mlp.linear_fc1.non_gated = True + elif change == "activation": + mlp.activation_func = torch.nn.functional.gelu + elif change == "fp8": + mlp.config.fp8 = "hybrid" + elif change == "tensor_parallel": + mlp.config.tensor_model_parallel_size = 2 + elif change == "bias": + mlp.config.add_bias_linear = True + elif change == "out_features": + mlp.linear_fc1.out_features = FFN + elif change == "other_layer": + layers[1].mlp = torch.nn.Linear(1, 1) + else: + mlp.linear_fc1 = mlp.linear_fc1.linear_fc1 + models = [_dense_model(layers)] * (2 if change == "chunks" else 1) + assert _dense_mlp_recompute_bytes_per_token(models) == 0 + + +@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 workspace below the gradient stage. + values = r._estimate_flat_forward(requests(rows, 16)) + 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. + per_row = r._recomputed_mixer_bytes_per_token() + 2 * HIDDEN * 2 + STAGE + assert cost.checkpoint_workspace >= rows * per_row + 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(_dense_rank(0), values) + assert plain.checkpoint_input_gradient == rows * LAYERS * HIDDEN * 2 + assert plain.checkpoint_workspace < cost.checkpoint_workspace + + +def test_dense_stage_is_not_used_beyond_cp2(monkeypatch): + r = _dense_rank() + values = r._estimate_flat_forward(requests(67, 4096)) + monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 4, 1)) + assert r._dense_recompute_stage_bytes() == 0 + cost = price(r, values) + assert cost.checkpoint_input_gradient == 67 * LAYERS * HIDDEN * 2 + + +def test_layout_floor_prices_the_dense_stage_per_layer_type(): + r = _dense_rank() + layers = [SimpleNamespace() for _ in range(4)] + layout = _GroupLayout( + attention_rows=(100, 80), gdn_rows=None, attention_retained=(0, 0) + ) + dense = r._layout_checkpoint_floor(layers, (None,), (None,), (layout,)) + plain = _dense_rank(0)._layout_checkpoint_floor(layers, (None,), (None,), (layout,)) + assert dense[0] == plain[0] == 100 * 4 * HIDDEN * 2 + # The busiest rank's recomputed layer adds its residual, norm and stage. + assert sum(dense) - sum(plain) == 100 * (2 * HIDDEN * 2 + STAGE) + + +def _no_grad_required(r, group_rows, *, topology=(1, 1, 2, 1), grad=False): + signature = _MemorySignature( + topology, (1, None), len(group_rows), (), grad, (grad,) * len(group_rows) + ) + return r._estimate_required_memory_bytes_from_values( + packed_tokens=40_000, + output_bytes=0, + signature=signature, + logical_tokens=40_000, + group_rows=tuple((rows, grad) for rows in group_rows), + ) + + +def test_no_grad_groups_price_only_the_largest_one(): + dense, plain = _dense_rank(), _dense_rank(0) + # One group: the traced per-packed-token floor, unchanged. + assert _no_grad_required(dense, (12_000,)) == _no_grad_required(plain, (12_000,)) + # Two sequential groups: only the larger group's share of the packed rows. + two = _no_grad_required(plain, (12_000, 8_000)) + assert _no_grad_required(dense, (12_000, 8_000)) == pytest.approx( + two * 12_000 / 20_000, rel=1e-3 + ) + + +@pytest.mark.parametrize("case", ["cp1", "cp4", "unsupported"]) +def test_no_grad_group_floor_needs_the_traced_shape(case): + dense = _dense_rank(0 if case == "unsupported" else STAGE) + kwargs: dict[str, Any] = { + "cp1": {"topology": (1, 1, 1, 1)}, + "cp4": {"topology": (1, 1, 4, 1)}, + "unsupported": {}, + }[case] + plain = _dense_rank(0) + assert _no_grad_required(dense, (12_000, 8_000), **kwargs) == _no_grad_required( + plain, (12_000, 8_000), **kwargs + ) 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, From b8bc47e0d3ad06f9b0ce8f7c12dd947542919eda Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:43:00 +0000 Subject: [PATCH 02/16] Charge dense no-grad transients and recompile residue; harden the gate Review follow-ups: - A no-grad group in a covered dense model is charged its traced 6F + 6H transient per row, including beside retained gradient boundaries in mixed waves, instead of 4H. - Multi-group no-grad waves price the largest group's own rows at that width, and at least its share of the per-token floor, instead of converting through the wave's average packed-to-row ratio. - The gradient stage adds one more FC1 triplet (6F) for the early-run recompile residue seen in q062, so cold waves need no learned profile. - Only at CP2, where it was traced. TE workspace growth is charged for dense too. - The gate requires the traced fused-norm FC1, the unfused SwiGLU config, no hooks or overrides on the decoder, layers, mixers or MLP besides ART's GDN wrappers, the trainer's hidden size, and LoRA ranks up to 256, rechecked per active slot. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 247 +++++++++---- src/art/trainer_rank/_planner_misses.py | 4 +- tests/unit/test_trainer_rank_dense_memory.py | 345 ++++++++++++++----- 3 files changed, 447 insertions(+), 149 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 6f205935c..68118eee8 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1630,34 +1630,53 @@ def _hybridep_buffer_bytes(capacity: int, ranks: int, hidden: int, experts: int) return tokens * (2 * hidden + 5 * experts + 4 * (hidden // 128)) -def _dense_mlp_recompute_bytes_per_token(model: Sequence[torch.nn.Module]) -> int: - """The dense MLP stage live at a recomputed layer's backward peak, per row. - - Qwen3.8-27B allocator traces at CP2 (dense, gated SwiGLU, LoRA on FC1 and - FC2): the peak sits in the recomputed layer's FC1 stage, holding the FC1 - base output, the LoRA gate/up output and their sum (2F each) plus one - F-wide tensor: 7F elements per row (238 KB measured, 244 KB priced at - F = 17,408). Every decoder layer must be this exact supported MLP; - otherwise 0 keeps the per-boundary gradient allowance. +# 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 LoRA rank intermediates. Every decoder layer, and ``slot_ref``'s + adapters, must match the traced execution (ART's own GDN layer and mixer + wrappers included); otherwise (0, 0) keeps today's allowances. """ if len(model) != 1: - return 0 + return 0, 0 try: decoder = _language_model(model[0]).decoder except (AttributeError, RuntimeError): - return 0 + return 0, 0 layers = getattr(decoder, "layers", None) if not layers or not all(hasattr(layer, "mlp") for layer in layers): - return 0 + return 0, 0 try: from megatron.core.extensions.transformer_engine import ( - TEColumnParallelLinear, TELayerNormColumnParallelLinear, TERowParallelLinear, ) from megatron.core.transformer.mlp import MLP from megatron.core.transformer.transformer_block import TransformerBlock + from art.megatron.gdn.operator import ( + _gdn_island_layer_forward, + _prefix_tree_forward, + ) from art.megatron.lora import ( LoRA, SelfAttentionLinearProjLoRA, @@ -1666,54 +1685,109 @@ def _dense_mlp_recompute_bytes_per_token(model: Sequence[torch.nn.Module]) -> in ) except ImportError: # Without the traced owner types nothing can match; keep the allowance. - return 0 - if type(decoder) is not TransformerBlock: - return 0 - stage = 0 + return 0, 0 + + def plain(module: Any, *wrappers: Any) -> bool: + """No hooks, and no forward override but ART's traced wrappers.""" + forward = vars(module).get("forward") + return ( + not module._forward_hooks + and not module._forward_pre_hooks + and ( + forward is None + or type(forward) is MethodType + and forward.__self__ is module + and forward.__func__ in wrappers + ) + ) + + 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, + "bias_activation_fusion": False, + "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) - base1 = getattr(fc1, "linear_fc1", 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, "lora", None), LoRA), (getattr(row, "linear_proj", None), TERowParallelLinear), - (getattr(fc1, "gate_lora", None), LoRA), - (getattr(fc1, "up_lora", None), LoRA), + *((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(base1) not in (TEColumnParallelLinear, TELayerNormColumnParallelLinear) + not isinstance(layer, torch.nn.Module) + or not plain(layer, _gdn_island_layer_forward) + or not isinstance(mixer, torch.nn.Module) + or not plain(mixer, _prefix_tree_forward) or any(type(site) is not cls for site, cls in sites) - or any( - "forward" in vars(site) - or cast(Any, site)._forward_hooks - or cast(Any, site)._forward_pre_hooks - for site in (base1, *(site for site, _ 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 getattr(config, "gated_linear_unit", None) is not True - or getattr(config, "params_dtype", None) is not torch.bfloat16 - or getattr(config, "add_bias_linear", None) is not False - or getattr(config, "sequence_parallel", None) is not False + 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, "cuda_graph_impl", "none") != "none" - or getattr(config, "tensor_model_parallel_size", None) != 1 - or getattr(config, "pipeline_model_parallel_size", None) != 1 + 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 - stage = max(stage, 7 * ffn * 2) - return stage + return 0, 0 + for adapter in adapters: + tensors = _slot_lora_tensors(adapter, 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 + rank = max(rank, int(a.shape[1])) + width, hidden = max(width, ffn), max(hidden, size) + # Each of three adapters keeps its rank-wide input product and gradient. + adapters = 6 * rank + return (7 * width + 6 * width + adapters) * 2, ( + 6 * width + 6 * hidden + adapters + ) * 2 def _moe_output_bytes_per_token( @@ -2177,10 +2251,15 @@ def memory_field(name: str, default: Any = None) -> Any: ) == 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 = ( - 0 + ( + 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) + 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( @@ -4369,7 +4448,8 @@ def _checkpoint_memory_floor( return self._layout_checkpoint_floor(decoder.layers, refs, routed, layouts) retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 - dense = self._dense_recompute_stage_bytes() 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; a # covered dense layer keeps its MLP stage. @@ -4389,12 +4469,12 @@ def _checkpoint_memory_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: workspace += self._te_workspace_growth_bytes() return retained, workspace @@ -4428,7 +4508,7 @@ 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() - dense = self._dense_recompute_stage_bytes() + dense, _ = self._dense_mlp_widths(refs) beside = ( 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe @@ -4469,7 +4549,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 @@ -4688,7 +4768,7 @@ def _checkpoint_input_gradient_bytes( CP2/EP1 and EP2/CP2), charge that gradient. Elsewhere keep one gradient per boundary: that allowance also covers dense MLP and other recompute work the floor does not price. A covered dense model (every layer the - traced gated MLP, ``_dense_recompute_stage_bytes``) also holds one: + 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) @@ -4696,7 +4776,9 @@ def _checkpoint_input_gradient_bytes( 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 self._dense_recompute_stage_bytes() or ( + if self._dense_mlp_widths( + tuple(ref for (_, grad), ref in zip(group_rows, refs, strict=True) if grad) + )[0] or ( self._checkpoint_moe_bytes_per_token() and all( self._moe_recompute_covered_for(ref) @@ -4707,17 +4789,37 @@ def _checkpoint_input_gradient_bytes( return gradient_rows * self._hidden_size * 2 return retained - def _dense_recompute_stage_bytes(self) -> int: - """The covered dense MLP stage per recomputed row, where it was traced. + 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 up to CP2: at higher CP a rank's remote attention stages may keep - more than the CP2 allowance, which the per-boundary gradient allowance - still covers there. + 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. """ stage = getattr(self, "_dense_recompute_bytes_per_token", 0) - if type(stage) is not int or stage <= 0 or self._topology_key()[2] > 2: - return 0 - return stage + 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 + ): + 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( @@ -6909,6 +7011,7 @@ def _fill_planner_snapshot( "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) @@ -8456,25 +8559,29 @@ def _estimate_required_memory_bytes_from_values( return output_bytes profiled = self._memory_profiles.get(signature) activation_factor = max(4, min(16, self._num_layers // 4 + 4)) - floor_tokens = packed_tokens - if ( - not signature.grad_enabled - and signature.topology[2] == 2 - and len(group_rows) > 1 - and self._dense_recompute_stage_bytes() - ): - # No-grad groups run one after another and keep nothing but their - # outputs (charged below), so only the largest group's transient - # is live. Qwen3.8-27B CP2 traces: 263 KB per busiest-rank row of - # one group, which this floor matches for a single group. - rows = [rows for rows, _ in group_rows] - floor_tokens = -(-packed_tokens * max(rows) // max(1, sum(rows))) static_compute = ( - floor_tokens + packed_tokens * self._hidden_size * 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. + rows = [rows for rows, _ in group_rows] + static_compute = 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 566cdc126..ec3b39d64 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -527,7 +527,9 @@ def report( "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"}) +_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]: diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index 65bb04aa3..a854940fe 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -1,25 +1,33 @@ """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 (7F per row) and one input gradient, and no-grad groups run one -after another. +FC1 stage and one input gradient, and no-grad groups run one after another. """ -from types import SimpleNamespace +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 _impl from art.trainer_rank._impl import ( _dense_mlp_recompute_bytes_per_token, _GroupLayout, _MemorySignature, ) -HIDDEN, LAYERS, FFN = 2048, 40, 5632 -STAGE = 7 * FFN * 2 +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): @@ -28,8 +36,15 @@ def _module(cls): 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() -> Any: - """The supported gated MLP, from the real owner types.""" + """The traced gated MLP, from the real owner types.""" pytest.importorskip("art.megatron.lora") from megatron.core.extensions.transformer_engine import ( TELayerNormColumnParallelLinear, @@ -41,32 +56,40 @@ def _dense_layer() -> Any: 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, - fp8=None, - fp4=None, + bias_activation_fusion=False, + 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 = _module(lora_module.LoRA) - fc1.up_lora = _module(lora_module.LoRA) + 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 = _module(lora_module.LoRA) + 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 = torch.nn.Module() + layer.self_attention = torch.nn.Module() layer.mlp = mlp return layer @@ -82,70 +105,190 @@ def _dense_model(layers: list[Any]) -> Any: return model -def _dense_rank(stage: int = STAGE): +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 test_supported_dense_mlp_prices_its_traced_stage(): +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(): + from art.megatron.gdn.operator import ( + _gdn_island_layer_forward, + _prefix_tree_forward, + ) + layers = [_dense_layer() for _ in range(3)] - assert _dense_mlp_recompute_bytes_per_token([_dense_model(layers)]) == STAGE + # ART's own GDN layer and mixer wrappers were part of the traced run. + layers[0].forward = MethodType(_gdn_island_layer_forward, layers[0]) + mixer = layers[0].self_attention + mixer.forward = MethodType(_prefix_tree_forward, mixer) + 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", [ - "hook", + "fc1_hook", + "row_hook", + "lora_hook", + "base_hook", + "layer_hook", + "layer_forward", + "mixer_hook", + "mixer_forward", + "no_mixer", + "decoder_hook", + "decoder_subclass", "non_gated", "activation", + "config_activation", + "clamp", + "linear_offset", + "fused_activation", + "te_activation", "fp8", + "fp4", + "dtype", "tensor_parallel", + "pipeline_parallel", + "sequence_parallel", + "cuda_graphs", "bias", "out_features", "other_layer", "unwrapped_fc1", + "unfused_norm", + "rank", "chunks", ], ) -def test_any_unsupported_layer_keeps_the_boundary_allowance(change): +def test_anything_but_the_traced_execution_keeps_the_allowance(change): + from megatron.core.extensions.transformer_engine import TEColumnParallelLinear + from megatron.core.transformer.transformer_block import TransformerBlock + + from art.megatron import lora as lora_module + layers = [_dense_layer() for _ in range(3)] - mlp = layers[1].mlp - if change == "hook": - mlp.linear_fc1.register_forward_hook(lambda *args: None) - elif change == "non_gated": - mlp.linear_fc1.non_gated = True - elif change == "activation": - mlp.activation_func = torch.nn.functional.gelu - elif change == "fp8": - mlp.config.fp8 = "hybrid" - elif change == "tensor_parallel": - mlp.config.tensor_model_parallel_size = 2 - elif change == "bias": - mlp.config.add_bias_linear = True - elif change == "out_features": - mlp.linear_fc1.out_features = FFN - elif change == "other_layer": - layers[1].mlp = torch.nn.Linear(1, 1) - else: - mlp.linear_fc1 = mlp.linear_fc1.linear_fc1 - models = [_dense_model(layers)] * (2 if change == "chunks" else 1) - assert _dense_mlp_recompute_bytes_per_token(models) == 0 + model = _dense_model(layers) + layer = layers[1] + mlp = layer.mlp + config = mlp.config + + def hook(module): + return lambda: module.register_forward_hook(lambda *args: None) + + class Block(TransformerBlock): + pass + + 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) + ), + "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), + "fused_activation": lambda: setattr(config, "bias_activation_fusion", True), + "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) + ), + "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 workspace below the gradient stage. + # 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. + # 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 + assert cost.checkpoint_workspace >= rows * per_row + r._te_workspace_growth_bytes() assert cost.required >= int( ( cost.checkpoint_retained @@ -155,66 +298,112 @@ def test_covered_dense_recompute_charges_one_gradient_and_its_stage(rows): * 1.1 ) # Without the traced stage, dense keeps one gradient per boundary. - plain = price(_dense_rank(0), values) + plain = price(_at_cp2(_dense_rank(0, 0)), values) assert plain.checkpoint_input_gradient == rows * LAYERS * HIDDEN * 2 - assert plain.checkpoint_workspace < cost.checkpoint_workspace -def test_dense_stage_is_not_used_beyond_cp2(monkeypatch): +def test_a_larger_no_grad_group_keeps_its_transient_beside_gradient_boundaries(): r = _dense_rank() - values = r._estimate_flat_forward(requests(67, 4096)) - monkeypatch.setattr(r, "_topology_key", lambda: (1, 1, 4, 1)) - assert r._dense_recompute_stage_bytes() == 0 - cost = price(r, values) - assert cost.checkpoint_input_gradient == 67 * LAYERS * HIDDEN * 2 + 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) -def test_layout_floor_prices_the_dense_stage_per_layer_type(): +@pytest.mark.parametrize("topology", [(1, 1, 1, 1), (1, 1, 4, 1)]) +def test_dense_widths_apply_only_at_cp2(topology): r = _dense_rank() - layers = [SimpleNamespace() for _ in range(4)] + 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("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=None, attention_retained=(0, 0) + 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,)) - plain = _dense_rank(0)._layout_checkpoint_floor(layers, (None,), (None,), (layout,)) - assert dense[0] == plain[0] == 100 * 4 * HIDDEN * 2 - # The busiest rank's recomputed layer adds its residual, norm and stage. - assert sum(dense) - sum(plain) == 100 * (2 * HIDDEN * 2 + STAGE) + 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=(1, 1, 2, 1), grad=False): +def _no_grad_required(r, group_rows, *, topology=CP2, packed=40_000): signature = _MemorySignature( - topology, (1, None), len(group_rows), (), grad, (grad,) * len(group_rows) + topology, (1, None), len(group_rows), (), False, (False,) * len(group_rows) ) return r._estimate_required_memory_bytes_from_values( - packed_tokens=40_000, + packed_tokens=packed, output_bytes=0, signature=signature, - logical_tokens=40_000, - group_rows=tuple((rows, grad) for rows in group_rows), + logical_tokens=packed, + group_rows=tuple((rows, False) for rows in group_rows), ) -def test_no_grad_groups_price_only_the_largest_one(): - dense, plain = _dense_rank(), _dense_rank(0) - # One group: the traced per-packed-token floor, unchanged. +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) + # One group: today's per-packed-token floor, unchanged. assert _no_grad_required(dense, (12_000,)) == _no_grad_required(plain, (12_000,)) - # Two sequential groups: only the larger group's share of the packed rows. - two = _no_grad_required(plain, (12_000, 8_000)) - assert _no_grad_required(dense, (12_000, 8_000)) == pytest.approx( - two * 12_000 / 20_000, rel=1e-3 - ) + 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, and at + # least its share of today's per-packed-token floor. + expected = max( + largest * NO_GRAD, -(-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 * 1.1) @pytest.mark.parametrize("case", ["cp1", "cp4", "unsupported"]) def test_no_grad_group_floor_needs_the_traced_shape(case): - dense = _dense_rank(0 if case == "unsupported" else STAGE) - kwargs: dict[str, Any] = { - "cp1": {"topology": (1, 1, 1, 1)}, - "cp4": {"topology": (1, 1, 4, 1)}, - "unsupported": {}, - }[case] - plain = _dense_rank(0) - assert _no_grad_required(dense, (12_000, 8_000), **kwargs) == _no_grad_required( - plain, (12_000, 8_000), **kwargs - ) + 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) From bf98ce860e5ae3a03b401023db42879eab0ea9ed Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:57:36 +0000 Subject: [PATCH 03/16] Decide dense eligibility once per wave; check wrapper delegates Round-2 review follow-ups: - The one-gradient decision uses the same slots as the floor, so an unsupported slot in any group keeps both allowances. - No-grad-only waves charge TE workspace growth too. - The gate requires Megatron's TransformerLayer, attention or GDN mixers whose class forward is the base one, and ART's GDN wrappers delegating to that class forward. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 58 ++++++---- tests/unit/test_trainer_rank_dense_memory.py | 109 ++++++++++++++++--- 2 files changed, 130 insertions(+), 37 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 68118eee8..a8e2d0f7d 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1670,8 +1670,11 @@ def _dense_mlp_recompute_bytes_per_token( 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 ( _gdn_island_layer_forward, @@ -1687,19 +1690,29 @@ def _dense_mlp_recompute_bytes_per_token( # Without the traced owner types nothing can match; keep the allowance. return 0, 0 - def plain(module: Any, *wrappers: Any) -> bool: - """No hooks, and no forward override but ART's traced wrappers.""" + 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) return ( - not module._forward_hooks - and not module._forward_pre_hooks - and ( - forward is None - or type(forward) is MethodType - and forward.__self__ is module - and forward.__func__ in wrappers - ) - ) + 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 + ) + + mixers = { + SelfAttention: SelfAttention.forward, + GatedDeltaNet: GatedDeltaNet.forward, + } if type(decoder) is not TransformerBlock or not plain(decoder): return 0, 0 @@ -1739,10 +1752,17 @@ def plain(module: Any, *wrappers: Any) -> bool: size = getattr(config, "hidden_size", None) mixer = getattr(layer, "self_attention", None) if ( - not isinstance(layer, torch.nn.Module) - or not plain(layer, _gdn_island_layer_forward) - or not isinstance(mixer, torch.nn.Module) - or not plain(mixer, _prefix_tree_forward) + type(layer) is not TransformerLayer + or not plain( + layer, _gdn_island_layer_forward, "_art_gdn_island_physical_forward" + ) + # Attention or GDN, with its base class's forward (Qwen subclasses + # keep it). + or not any( + isinstance(mixer, base) and type(mixer).forward is forward + for base, forward in mixers.items() + ) + 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 @@ -4474,7 +4494,7 @@ def _checkpoint_memory_floor( group_rows, refs, routed, strict=True ) ) - if moe or dense: + if moe or dense or no_grad: workspace += self._te_workspace_growth_bytes() return retained, workspace @@ -4776,9 +4796,9 @@ def _checkpoint_input_gradient_bytes( 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 self._dense_mlp_widths( - tuple(ref for (_, grad), ref in zip(group_rows, refs, strict=True) if grad) - )[0] or ( + # 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._checkpoint_moe_bytes_per_token() and all( self._moe_recompute_covered_for(ref) diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index a854940fe..3c19cd7d6 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -43,14 +43,17 @@ def _adapter(lora_module, inputs: int, outputs: int, rank: int = RANK): return lora -def _dense_layer() -> Any: +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 @@ -88,12 +91,27 @@ def _dense_layer() -> Any: fc2.row_parallel_lora = row mlp.linear_fc1 = fc1 mlp.linear_fc2 = fc2 - layer = torch.nn.Module() - layer.self_attention = torch.nn.Module() + 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, + ) + + layer._art_gdn_island_physical_forward = 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 @@ -121,16 +139,10 @@ def _at_cp2(r): def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): - from art.megatron.gdn.operator import ( - _gdn_island_layer_forward, - _prefix_tree_forward, - ) - - layers = [_dense_layer() for _ in range(3)] - # ART's own GDN layer and mixer wrappers were part of the traced run. - layers[0].forward = MethodType(_gdn_island_layer_forward, layers[0]) - mixer = layers[0].self_attention - mixer.forward = MethodType(_prefix_tree_forward, mixer) + # 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) == ( @@ -150,6 +162,10 @@ def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): "base_hook", "layer_hook", "layer_forward", + "layer_delegate", + "layer_type", + "mixer_delegate", + "mixer_class_forward", "mixer_hook", "mixer_forward", "no_mixer", @@ -180,13 +196,18 @@ def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): ) 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 - layers = [_dense_layer() for _ in range(3)] + 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 @@ -196,6 +217,15 @@ def hook(module): 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) + edits = { "fc1_hook": hook(mlp.linear_fc1), "row_hook": hook(mlp.linear_fc2.row_parallel_lora), @@ -205,6 +235,18 @@ class Block(TransformerBlock): "layer_forward": lambda: setattr( layer, "forward", MethodType(lambda self, *a: None, layer) ), + "layer_delegate": lambda: setattr( + layer, "_art_gdn_island_physical_forward", 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 @@ -379,14 +421,16 @@ 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, and at - # least its share of today's per-packed-token floor. + # 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, -(-packed * per_token * largest // sum(groups)) + 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( @@ -394,7 +438,7 @@ def test_no_grad_groups_price_the_largest_groups_own_rows(): ) # 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 * 1.1) + assert _no_grad_required(wide, (12_000, 8_000)) == int((12_000 * 10**6 + te) * 1.1) @pytest.mark.parametrize("case", ["cp1", "cp4", "unsupported"]) @@ -407,3 +451,32 @@ def test_no_grad_group_floor_needs_the_traced_shape(case): 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 From 05188739f098f38d1c18e17181a48686cc0a406d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 09:02:49 +0000 Subject: [PATCH 04/16] Accept Dynamo-compiled GDN layer delegates in the dense gate Training compile replaces each layer's _art_gdn_island_physical_forward with torch.compile's wrapper, so the gate rejected every layer of the compiled model it was traced on. Judge the callable Dynamo wraps. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 4 ++++ tests/unit/test_trainer_rank_dense_memory.py | 7 ++++++- 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index a8e2d0f7d..03a8f5617 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1699,6 +1699,10 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: 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 diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index 3c19cd7d6..c5bf16ce5 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -104,7 +104,8 @@ def _wrap_like_art(layer: Any) -> None: _prefix_tree_forward, ) - layer._art_gdn_island_physical_forward = layer.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": @@ -163,6 +164,7 @@ def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): "layer_hook", "layer_forward", "layer_delegate", + "compiled_custom_delegate", "layer_type", "mixer_delegate", "mixer_class_forward", @@ -238,6 +240,9 @@ def forward(self, *args, **kwargs): # A class-level override. "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, From 7ff82bcb0a91b2f8dd910dc76856b270b16132df Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 09:13:06 +0000 Subject: [PATCH 05/16] Match the dense gate to the traced Qwen execution The traced Qwen3.8-27B run uses Megatron Bridge's Qwen3VLSelfAttention (which overrides forward) and Bridge's fused SwiGLU (bias_activation_fusion=True), so the gate rejected every layer of it. Accept exactly the traced mixer types and the fused activation. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 26 +++++++++++--------- tests/unit/test_trainer_rank_dense_memory.py | 21 +++++++++++++--- 2 files changed, 33 insertions(+), 14 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 03a8f5617..8e86d3c07 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1713,10 +1713,17 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: and inner.__func__ is type(module).forward ) - mixers = { - SelfAttention: SelfAttention.forward, - GatedDeltaNet: GatedDeltaNet.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 @@ -1725,7 +1732,9 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: "params_dtype": torch.bfloat16, "add_bias_linear": False, "sequence_parallel": False, - "bias_activation_fusion": 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", @@ -1760,12 +1769,7 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: or not plain( layer, _gdn_island_layer_forward, "_art_gdn_island_physical_forward" ) - # Attention or GDN, with its base class's forward (Qwen subclasses - # keep it). - or not any( - isinstance(mixer, base) and type(mixer).forward is forward - for base, forward in mixers.items() - ) + 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) diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index c5bf16ce5..ca9f66529 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -65,7 +65,7 @@ def _dense_layer(gdn: bool = False) -> Any: params_dtype=torch.bfloat16, add_bias_linear=False, sequence_parallel=False, - bias_activation_fusion=False, + bias_activation_fusion=True, # Bridge's Qwen3.5 providers fuse SwiGLU. use_te_activation_func=False, cpu_offloading=False, cuda_graph_impl="none", @@ -178,7 +178,7 @@ def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): "config_activation", "clamp", "linear_offset", - "fused_activation", + "unfused_activation", "te_activation", "fp8", "fp4", @@ -266,7 +266,7 @@ def forward(self, *args, **kwargs): # A class-level override. ), "clamp": lambda: setattr(config, "activation_func_clamp_value", 7.0), "linear_offset": lambda: setattr(config, "glu_linear_offset", 1.0), - "fused_activation": lambda: setattr(config, "bias_activation_fusion", True), + "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"), @@ -485,3 +485,18 @@ def test_no_grad_only_waves_charge_te_workspace_growth(): _, 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, + ) From da1294dfa400ff2297a2848a03bf8ec1f52c2794 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 10:01:35 +0000 Subject: [PATCH 06/16] Check everything a dense layer runs before pricing its traced stage The gate now walks every module under each decoder layer, not just the MLP: - no hooks or instance forwards except ART's GDN layer, mixer and empty-safe norm wrappers, which must delegate to the class forward; - every adapter, including the mixer's, is an exact LoRA with no selector override, within the rank limit; - every adapter's rank term is priced. Dense widths now require the checkpoint floor's own decoder conditions, so the multi-group no-grad discount always keeps its TE workspace growth. The split planner's optimistic lower bound drops the largest group's share of the per-token floor. That share shrinks as another group's rows grow, so it could exceed the exact cost and prune a feasible split. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 132 +++++++++++----- tests/unit/test_trainer_rank_dense_memory.py | 154 ++++++++++++++++++- 2 files changed, 245 insertions(+), 41 deletions(-) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 8e86d3c07..465f2b67f 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1652,9 +1652,10 @@ def _dense_mlp_recompute_bytes_per_token( - 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 LoRA rank intermediates. Every decoder layer, and ``slot_ref``'s - adapters, must match the traced execution (ART's own GDN layer and mixer - wrappers included); otherwise (0, 0) keeps today's allowances. + 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 @@ -1677,6 +1678,7 @@ def _dense_mlp_recompute_bytes_per_token( 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, ) @@ -1793,8 +1795,30 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: or getattr(config, "glu_linear_offset", 0.0) != 0.0 ): return 0, 0 - for adapter in adapters: - tensors = _slot_lora_tensors(adapter, slot_ref) + # 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 @@ -1809,10 +1833,12 @@ def plain(module: Any, wrapper: Any = None, delegate: str = "") -> bool: or not 0 < a.shape[1] <= _DENSE_LORA_RANK_LIMIT ): return 0, 0 - rank = max(rank, int(a.shape[1])) + layer_rank += int(a.shape[1]) + rank = max(rank, layer_rank) width, hidden = max(width, ffn), max(hidden, size) - # Each of three adapters keeps its rank-wide input product and gradient. - adapters = 6 * rank + # 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 @@ -3936,6 +3962,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 ( @@ -4393,46 +4420,28 @@ def _moe_workspace_bytes( else routed * coefficient ) - def _checkpoint_memory_floor( - self, - group_rows: tuple[tuple[int, bool], ...], - slot_refs: tuple["LoRASlotRef | None", ...] | None = None, - routed_rows: tuple[int, ...] | None = None, - layouts: tuple[_GroupLayout, ...] | None = None, - ) -> tuple[int, int]: - """Conservative saved-boundary charge and one recomputed layer's workspace. + def _checkpoint_floor_decoder(self) -> 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. + Every local layer recomputed full/uniform/1 in BF16 at TP1/PP1, with + no custom checkpointed forward. The dense widths require it too, so a + discount never outlives the floor that carries its TE growth. """ - 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) expected = { @@ -4469,7 +4478,38 @@ 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, + 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. + """ + gradient_rows = sum(rows for rows, grad in group_rows if grad) + decoder = self._checkpoint_floor_decoder() 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 layouts is not None and all(grad for _, grad in group_rows): @@ -4826,7 +4866,8 @@ def _dense_mlp_widths( 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. + 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) @@ -4836,6 +4877,7 @@ def _dense_mlp_widths( 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 (): @@ -4882,7 +4924,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, @@ -4897,6 +4941,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, slot_refs, group_routed_rows, group_layouts @@ -8582,6 +8627,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 @@ -8605,10 +8651,16 @@ def _estimate_required_memory_bytes_from_values( # (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( - max(rows) * no_grad, - -(-static_compute * max(rows) // max(1, sum(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 diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index ca9f66529..0a5f1373e 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -14,8 +14,9 @@ from test_trainer_rank_moe_memory import layer # noqa: F401 import torch -from art.trainer_rank import _impl +from art.trainer_rank import ForwardInput, _impl from art.trainer_rank._impl import ( + Unset, _dense_mlp_recompute_bytes_per_token, _GroupLayout, _MemorySignature, @@ -193,6 +194,12 @@ def test_traced_dense_mlp_prices_its_stage_and_no_grad_transient(): "unwrapped_fc1", "unfused_norm", "rank", + "mixer_adapter_rank", + "mixer_child_hook", + "mixer_child_forward", + "norm_delegate", + "slot_selector", + "active_selector", "chunks", ], ) @@ -203,6 +210,7 @@ def test_anything_but_the_traced_execution_keeps_the_allowance(change): 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: @@ -228,6 +236,19 @@ def forward(self, *args, **kwargs): # A class-level override. 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), @@ -285,6 +306,23 @@ def forward(self, *args, **kwargs): # A class-level override. "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]() @@ -500,3 +538,117 @@ def test_the_traced_qwen_attention_mixer_is_accepted(): 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] From 865eb6a0e15af9d0f850f97c4fac07db4d1b82e9 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:39:32 +0000 Subject: [PATCH 07/16] Test that dense widths stay at TP1 where a sequence-parallel floor prices Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_trainer_rank_dense_memory.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/tests/unit/test_trainer_rank_dense_memory.py b/tests/unit/test_trainer_rank_dense_memory.py index 0a5f1373e..e4e59c0c0 100644 --- a/tests/unit/test_trainer_rank_dense_memory.py +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -406,6 +406,25 @@ def test_dense_widths_apply_only_at_cp2(topology): 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()) From 98eec50d5256bd5cc93e0ad05cd5e83c23a35d48 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 06:46:47 +0000 Subject: [PATCH 08/16] Price dense recompute by its traced stage and one input gradient, on the extracted layout Port #986 (through fc6955b4e) onto the extracted planner modules: a dense model whose every layer is the traced gated MLP prices that MLP stage beside the recomputed mixer at CP2 (TP1), one input gradient, and no-grad groups at the traced transient width. The checkpoint floor's eligibility moves to _checkpoint_floor_decoder. Grouped planner-miss replay freezes the plan's dense widths (facts version 5). Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/prek.yml | 4 +- src/art/trainer_rank/_impl.py | 233 +++++- src/art/trainer_rank/_memory.py | 199 ++++-- src/art/trainer_rank/_micro_batch_planner.py | 5 +- src/art/trainer_rank/_planner_replay.py | 33 +- tests/unit/test_grouped_planner_replay.py | 78 +- tests/unit/test_trainer_rank_dense_memory.py | 707 +++++++++++++++++++ 7 files changed, 1199 insertions(+), 60 deletions(-) create mode 100644 tests/unit/test_trainer_rank_dense_memory.py diff --git a/.github/workflows/prek.yml b/.github/workflows/prek.yml index f4506e8af..9113bdc5d 100644 --- a/.github/workflows/prek.yml +++ b/.github/workflows/prek.yml @@ -247,6 +247,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 \ @@ -299,4 +300,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 7aef00465..5e0c80524 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -26,7 +26,7 @@ import threading import time import traceback -from types import TracebackType +from types import MethodType, TracebackType from typing import ( TYPE_CHECKING, Any, @@ -1675,6 +1675,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, @@ -2141,6 +2355,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, @@ -3059,7 +3285,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).""" checkpoint_memory = self._checkpoint_memory_floor( group_rows, slot_refs, @@ -3082,6 +3310,7 @@ def _subforward_cost( retained_tokens=retained_tokens, include_checkpoint_input_gradient=False, checkpoint_memory=checkpoint_memory, + lower_bound=lower_bound, ) checkpoint_retained, checkpoint_workspace = checkpoint_memory retained = self._retained_memory_bytes( @@ -4985,6 +5214,8 @@ def _cp_group_model_tokens( _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_floor_decoder = _memory._checkpoint_floor_decoder + _dense_mlp_widths = _memory._dense_mlp_widths _checkpoint_memory_floor = _memory._checkpoint_memory_floor _gradient_slots = staticmethod(_memory._gradient_slots) _pending_adapter_gradient_bytes = _memory._pending_adapter_gradient_bytes diff --git a/src/art/trainer_rank/_memory.py b/src/art/trainer_rank/_memory.py index b0d0da39b..8871882ac 100644 --- a/src/art/trainer_rank/_memory.py +++ b/src/art/trainer_rank/_memory.py @@ -534,39 +534,32 @@ def _moe_workspace_bytes( ) -def _checkpoint_layers( - self: TrainerRank, - group_rows: tuple[tuple[int, bool], ...], -) -> int: - """Conservative saved-boundary charge and one disjoint MoE workspace. - - 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. - No-grad groups also keep decoder input, current layer input, its MLP - residual and norm output across the MoE stage. Count these four row - tensors separately from returned outputs, allowing storage aliases. - This is not a bound for custom preprocessing, attention, or all 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. +def _checkpoint_floor_decoder( + self: TrainerRank, *, sequence_parallel: bool = False +) -> Any | None: + """The decoder whose saved boundaries the checkpoint floor prices, or None. + + 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. """ - if not group_rows or len(self.runtime.model) != 1: - return 0 + if len(self.runtime.model) != 1: + return None try: decoder = _impl._language_model(self.runtime.model[0]).decoder except (AttributeError, RuntimeError): - return 0 + return None try: from megatron.core.transformer.transformer_block import TransformerBlock except ModuleNotFoundError as error: if error.name != "megatron": raise - return 0 + return None if type(decoder) is not TransformerBlock: - return 0 + return None config = decoder.config layers = len(decoder.layers) _, tp, cp, pp = self._topology_key() @@ -575,7 +568,7 @@ def _checkpoint_layers( "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", @@ -590,7 +583,11 @@ def _checkpoint_layers( or self._param_dtype_size != 2 or next(self.runtime.model[0].parameters()).dtype is not _impl.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 @@ -605,8 +602,33 @@ def _checkpoint_layers( or getattr(decoder, "_forward_hooks", None) or getattr(decoder, "_forward_pre_hooks", None) ): + return None + return decoder + + +def _checkpoint_layers( + self: TrainerRank, + group_rows: tuple[tuple[int, bool], ...], +) -> int: + """Conservative saved-boundary charge and one disjoint MoE workspace. + + 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. + No-grad groups also keep decoder input, current layer input, its MLP + residual and norm output across the MoE stage. Count these four row + tensors separately from returned outputs, allowing storage aliases. + This is not a bound for custom preprocessing, attention, or all 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. + """ + if not group_rows: return 0 - return layers + _, tp, _, _ = self._topology_key() + decoder = self._checkpoint_floor_decoder(sequence_parallel=tp > 1) + return 0 if decoder is None else len(decoder.layers) def _checkpoint_memory_floor( @@ -713,13 +735,18 @@ def _checkpoint_floor_from_facts( # Boundaries on the busiest rank's rows and one recomputed layer. retained = gradient_rows * layers * self._hidden_size * 2 moe = self._checkpoint_moe_bytes_per_token() if gradient_rows else 0 + 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 @@ -729,14 +756,49 @@ def _checkpoint_floor_from_facts( 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 +def _dense_mlp_widths( + self: TrainerRank, 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 = _impl._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 _gradient_slots( group_rows: Sequence[tuple[int, bool]], slot_refs: Sequence[LoRASlotRef | None] | None, @@ -1064,8 +1126,11 @@ 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. ``retained`` is the - floor's boundary charge where the caller already has it. + 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`` is the floor's boundary charge where the caller already has + it. """ if retained is None: retained, _ = self._checkpoint_memory_floor(group_rows) @@ -1080,16 +1145,26 @@ def _checkpoint_gradient_covered( self: TrainerRank, group_rows: tuple[tuple[int, bool], ...], slot_refs: tuple["LoRASlotRef | None", ...] | None, + *, + dense: bool = True, ) -> bool: - """Whether the traced MoE stage covers every gradient group's recompute.""" + """Whether a traced stage covers every gradient group's recompute. + + The MoE stage, or with ``dense`` the covered dense MLP. + """ refs = (None,) * len(group_rows) if slot_refs is None else slot_refs + # The same slots as the floor: if any group's slot falls back there, + # the per-boundary allowance must stay here too. 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 + (dense and self._dense_mlp_widths(refs)[0]) + or ( + 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 + ) ) ) @@ -1114,15 +1189,16 @@ def _checkpoint_head_stage_bytes( 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``. - On per-rank CP layouts, each rank's boundaries are within the floor's, - so the largest rank's adapter term bounds each. + waves) show these terms at the head's peak. Elsewhere None, a covered + dense model included (untraced here): 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``. On + per-rank CP layouts, each rank's boundaries are within the floor's, so + the largest rank's adapter term bounds each. """ if not head_workspace_bytes or not self._checkpoint_gradient_covered( - group_rows, slot_refs + group_rows, slot_refs, dense=False ): return None rows = sum(rows for rows, grad in group_rows if grad) @@ -1278,7 +1354,14 @@ def _layout_checkpoint_rank_floors( attention_inputs = len(inputs) - gdn_inputs widths = self._recomputed_mixer_widths(stage_buffers=False) moe = self._checkpoint_moe_bytes_per_token() - beside = 2 * hidden + self._moe_checkpoint_state_bytes_per_token() if moe else 0 + 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 + ) floors: list[tuple[int, int]] = [] for rank in range(len(layouts[0].attention_rows)): retained = workspace = 0 @@ -1306,7 +1389,7 @@ def _layout_checkpoint_rank_floors( ) workspace = max(workspace, stage) floors.append((retained, workspace)) - growth = self._te_workspace_growth_bytes() if moe else 0 + growth = self._te_workspace_growth_bytes() if moe or dense else 0 return tuple((retained, workspace + growth) for retained, workspace in floors) @@ -2017,6 +2100,7 @@ def _estimate_required_memory_bytes_from_values( retained_tokens: int | None = None, include_checkpoint_input_gradient: bool = True, checkpoint_memory: tuple[int, int] | None = None, + lower_bound: bool = False, ) -> int: if packed_tokens <= 0: return output_bytes @@ -2025,6 +2109,29 @@ def _estimate_required_memory_bytes_from_values( static_compute = ( packed_tokens * self._hidden_size * 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/_micro_batch_planner.py b/src/art/trainer_rank/_micro_batch_planner.py index aa98f07d2..4220c5d1b 100644 --- a/src/art/trainer_rank/_micro_batch_planner.py +++ b/src/art/trainer_rank/_micro_batch_planner.py @@ -505,6 +505,7 @@ def _split_chunk_lower_cost( # admission cost. retained_tokens=(packed_tokens + signature.topology[2] - 1) // signature.topology[2], + lower_bound=True, ) for traced in _impl._traced_states(head_traced) ), @@ -601,7 +602,9 @@ def _plan_head_backward_traced(self: TrainerRank, plan: _FlatForwardPlan) -> boo 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) + self._plan_group_rows(plan), + tuple(g.slot_ref for g in plan.groups), + dense=False, ) ): object.__setattr__(plan, "_head_staged", True) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index e5a7b7a04..fc6f0e4a5 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -147,6 +147,8 @@ def capture(rank: Any, plan: Any) -> dict[str, Any]: "_layout_adapter_gradient_bytes", "_checkpoint_adapter_gradient_extra", "_recomputed_mixer_widths", + "_checkpoint_floor_decoder", + "_dense_mlp_widths", # Producers of recorded arguments replay checks against the facts. "_plan_group_routed_rows", "_plan_head_backward_traced", @@ -342,8 +344,14 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: if len(inputs) > MAX_LAYERS: raise ValueError("runtime_shape_inventory_over_limit") reserve(8 * len(inputs)) + # The covered dense stage for this plan's slots, and for no named slot + # (the layout-free floor behind the input-gradient allowance). + dense = [ + [int(v) for v in rank._dense_mlp_widths(refs)] + for refs in (tuple(g.slot_ref for g in plan.groups), None) + ] facts = { - "version": 4, + "version": 5, "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, @@ -353,6 +361,8 @@ def terms(checkpoint_grad: bool, ref: Any) -> list[Any]: for name in _RECOMPUTE_READERS }, "layer_gdn_inputs": list(inputs), + "dense_widths": dense[0], + "dense_base_widths": dense[1], "head_vocabulary": vocabulary, "head_target_backward": target_backward, "groups": groups, @@ -385,12 +395,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: "backward_row_state_bytes", "triton_min_rows", "layer_gdn_inputs", + "dense_widths", + "dense_base_widths", "head_vocabulary", "head_target_backward", "groups", }, ) - if type(facts["version"]) is not int or facts["version"] != 4: + if type(facts["version"]) is not int or facts["version"] != 5: raise ValueError("unsupported runtime facts version") for key in ( "checkpoint_layers", @@ -398,6 +410,11 @@ def integer(value: Any, *, minimum: int = 0) -> None: "head_vocabulary", ): integer(facts[key]) + for key in ("dense_widths", "dense_base_widths"): + if type(facts[key]) is not list or len(facts[key]) != 2: + raise ValueError("invalid dense stage facts") + for value in facts[key]: + integer(value) for key in _RECOMPUTE_READERS: if not facts["checkpoint_layers"]: if facts[key] is not None: @@ -763,6 +780,18 @@ def _checkpoint_memory_floor( layouts, ) + def _dense_mlp_widths(self, slot_refs: Any = None) -> tuple[int, int]: + if self._facts is None: + return _memory._dense_mlp_widths(self, slot_refs) + refs = tuple(slot_refs or ()) + if all(ref is None for ref in refs): + stage, no_grad = self._facts["dense_base_widths"] + elif refs == tuple(range(len(self._facts["groups"]))): + stage, no_grad = self._facts["dense_widths"] + else: + raise ValueError("replayed dense widths are frozen per plan") + return stage, no_grad + def _layer_gdn_inputs(self) -> tuple[bool, ...]: if self._facts is None: return _memory._layer_gdn_inputs(self) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 80e799ef5..6cdba1dba 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -409,20 +409,14 @@ def test_staged_head_needs_recorded_cp2_target_backward(tmp_path): 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}, - ) +def cp2(monkeypatch, rank): + """Plan and price ``rank`` at CP2 through the stock topology reader (as + capture requires), over a CPU CP2 topology and planning config.""" 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, @@ -436,6 +430,16 @@ def test_staged_cp2_head_is_replayed(monkeypatch, tmp_path): build_gdn_execution_spec=True, context_parallel_workload_profile=lambda provider: None, ) + return rank + + +def test_staged_cp2_head_is_replayed(monkeypatch, tmp_path): + monkeypatch.setattr( + tr, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + rank = cp2(monkeypatch, head_rank()) 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 @@ -447,6 +451,62 @@ def test_staged_cp2_head_is_replayed(monkeypatch, tmp_path): assert result["estimates"][0]["required_bytes"] == costs[0].required +def test_dense_stage_is_replayed_and_frozen(monkeypatch, tmp_path): + from test_trainer_rank_dense_memory import NO_GRAD, STAGE, _dense_rank + + rank = cp2(monkeypatch, _dense_rank()) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward([request(512, grad=True)]) + report, costs = emitted(rank, plan, tmp_path) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["dense_widths"] == facts["dense_base_widths"] == [STAGE, NO_GRAD] + actual = reports.replay(report) + assert actual["aggregate"]["matches"] + assert actual["estimates"][0]["required_bytes"] == costs[0].required + # The live constructor widths cannot change the replayed answer. + rank._dense_recompute_bytes_per_token = rank._dense_no_grad_bytes_per_token = 0 + assert reports.replay(report) == actual + changed = deepcopy(report) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["dense_widths"][ + 0 + ] += 10**6 + result = reports.replay(changed)["estimates"][0] + assert not result["matches"] + assert result["required_bytes"] > actual["estimates"][0]["required_bytes"] + + +@pytest.mark.parametrize("change", ["length", "value"]) +def test_dense_fact_validation_rejects_forged_input(change, monkeypatch, tmp_path): + from test_trainer_rank_dense_memory import _dense_rank + + rank = cp2(monkeypatch, _dense_rank()) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(512, grad=True)]), tmp_path + ) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + if change == "length": + facts["dense_widths"].append(0) + message = "invalid dense stage facts" + else: + facts["dense_base_widths"][1] = 1.5 + message = "invalid runtime dimension" + with pytest.raises(ValueError, match=message): + reports.replay(report) + + +@pytest.mark.parametrize("name", ["_dense_mlp_widths", "_checkpoint_floor_decoder"]) +def test_custom_dense_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, **kwargs: original(*args, **kwargs)) + with pytest.raises(ValueError, match="custom_runtime_estimator"): + _planner_replay.capture(rank, plan) + + @pytest.mark.parametrize("change", ["omit_default", "omit_required", "extra"]) def test_recorded_geometry_fields_follow_its_schema(change, tmp_path): rank = head_rank() 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..6320e0082 --- /dev/null +++ b/tests/unit/test_trainer_rank_dense_memory.py @@ -0,0 +1,707 @@ +"""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 ( + _COLD_RECOMPUTE_TRANSIENT_BYTES as COLD, +) +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_covered_dense_head_stays_in_the_decoder_stage(): + r = _dense_rank() + n, out, signature, groups, _ = r._estimate_flat_forward(requests(67, 16)) + _at_cp2(r) + # Staged head pricing was traced on the MoE model only. + head = 10**9 + cost = price(r, (n, out, signature, groups, head)) + gradient = cost.checkpoint_input_gradient + assert gradient == 67 * HIDDEN * 2 + assert r._checkpoint_head_stage_bytes(head, gradient, groups, None) is None + workspace = r._checkpoint_memory_floor(groups)[1] + assert cost.checkpoint_workspace == max(workspace, head) + COLD + assert cost.required == int( + (cost.checkpoint_retained + max(workspace, head) + COLD + gradient) * 1.1 + ) + + +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, + ) + # Each layer's saved input layout (``_layer_gdn_inputs``). + inputs = tuple(bool(gdn and index) for index in range(4)) + for x in (r, plain): + x._layer_gdn_inputs = lambda: inputs + layout = _GroupLayout( + attention_rows=(100, 80), + gdn_rows=(120, 60) if gdn else None, + attention_retained=(0, 0), + ) + dense = r._layout_checkpoint_floor((None,), (None,), (layout,)) + base = plain._layout_checkpoint_floor((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] + + +def test_covered_dense_plans_are_never_marked_staged(monkeypatch): + from art.trainer_rank import _impl + + monkeypatch.setattr( + _impl, + "_TRITON_STATS_STATE", + {"succeeded": {"local_logsumexp_stats"}, "failed": False}, + ) + r = _dense_rank() + monkeypatch.setattr(r, "_plan_head_workspace_bytes", lambda plan: 10**9) + monkeypatch.setattr(r, "_head_backward_traced", lambda *a, **k: True) + plan = r._plan_flat_forward(requests(67, 16)) + _at_cp2(r) + # Eligible, but the dense head keeps the unstaged price: not strict. + assert r._plan_head_backward_traced(plan) is True + assert getattr(plan, "_head_staged", False) is False From 78f3073d3f9cb7302b061b71eba28ff3b9054b17 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 07:29:13 +0000 Subject: [PATCH 09/16] Refuse dense widths live pricing cannot produce A dense pair is both zero or both positive; a plan's pair is zero or at least the constructor's; any nonzero pair needs a checkpointed decoder and the TP1/CP2/PP1 topology. Also replay dense widths on per-rank CP layouts in tests. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 12 ++++++ tests/unit/test_grouped_planner_replay.py | 45 ++++++++++++++++++++--- 2 files changed, 52 insertions(+), 5 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 8b8159130..ea8293399 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -415,6 +415,15 @@ def integer(value: Any, *, minimum: int = 0) -> None: raise ValueError("invalid dense stage facts") for value in facts[key]: integer(value) + # Live widths are both zero or both positive; a plan's named slots only + # raise the constructor's, and only a checkpointed decoder has any. + plan, base = facts["dense_widths"], facts["dense_base_widths"] + if ( + any(all(pair) != any(pair) for pair in (plan, base)) + or (any(plan) and not (all(base) and plan[0] >= base[0] and plan[1] >= base[1])) + or (any(base) and not facts["checkpoint_layers"]) + ): + raise ValueError("invalid dense stage facts") for key in _RECOMPUTE_READERS: if not facts["checkpoint_layers"]: if facts[key] is not None: @@ -852,6 +861,9 @@ def runtime_arguments( and facts["head_vocabulary"] ): raise ValueError("head backward staging disagrees with recorded facts") + # Live dense widths exist only at TP1/CP2/PP1 (_dense_mlp_widths). + if any(facts["dense_base_widths"]) and self._topology_key()[1:] != (1, 2, 1): + raise ValueError("dense stage facts disagree with the recorded topology") layouts = None if groups[0]["layout"] is not None: # Live layout pricing is CP2 at TP1/PP1 only (_layout_pricing_supported). diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 53bacb237..8868904ac 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -463,9 +463,6 @@ def test_dense_stage_is_replayed_and_frozen(monkeypatch, tmp_path): actual = reports.replay(report) assert actual["aggregate"]["matches"] assert actual["estimates"][0]["required_bytes"] == costs[0].required - # The live constructor widths cannot change the replayed answer. - rank._dense_recompute_bytes_per_token = rank._dense_no_grad_bytes_per_token = 0 - assert reports.replay(report) == actual changed = deepcopy(report) changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["dense_widths"][ 0 @@ -475,7 +472,32 @@ def test_dense_stage_is_replayed_and_frozen(monkeypatch, tmp_path): assert result["required_bytes"] > actual["estimates"][0]["required_bytes"] -@pytest.mark.parametrize("change", ["length", "value"]) +def test_dense_stage_on_cp_layouts_is_replayed(monkeypatch, tmp_path): + from test_trainer_rank_dense_memory import NO_GRAD, STAGE, _dense_rank + from test_trainer_rank_layout_memory import _requests, art_cp + + rank = art_cp(_dense_rank(), monkeypatch) + del rank._topology_key # the stock reader, over the fixture's CP2 topology + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + plan = rank._plan_flat_forward(_requests()) + assert rank._plan_group_layouts(plan) is not None + report, costs = emitted(rank, plan, tmp_path) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + assert facts["groups"][0]["layout"] is not None + assert facts["dense_widths"] == [STAGE, NO_GRAD] + actual = reports.replay(report) + assert actual["aggregate"]["matches"] + assert actual["estimates"][0]["required_bytes"] == costs[0].required + changed = deepcopy(report) + changed["replay"]["memory_replay"]["estimates"][0]["runtime_facts"]["dense_widths"][ + 0 + ] += 10**6 + assert not reports.replay(changed)["estimates"][0]["matches"] + + +@pytest.mark.parametrize( + "change", ["length", "value", "half_zero", "below_base", "no_base", "topology"] +) def test_dense_fact_validation_rejects_forged_input(change, monkeypatch, tmp_path): from test_trainer_rank_dense_memory import _dense_rank @@ -488,9 +510,22 @@ def test_dense_fact_validation_rejects_forged_input(change, monkeypatch, tmp_pat if change == "length": facts["dense_widths"].append(0) message = "invalid dense stage facts" - else: + elif change == "value": facts["dense_base_widths"][1] = 1.5 message = "invalid runtime dimension" + elif change == "half_zero": + facts["dense_widths"][1] = 0 + message = "invalid dense stage facts" + elif change == "below_base": + facts["dense_widths"][0] = facts["dense_base_widths"][0] - 1 + message = "invalid dense stage facts" + elif change == "no_base": + facts["dense_base_widths"] = [0, 0] + message = "invalid dense stage facts" + else: + # Dense widths on a report whose recorded topology is not CP2. + report["replay"]["memory_replay"]["rank"]["topology"] = [1, 1, 1, 1] + message = "dense stage facts disagree" with pytest.raises(ValueError, match=message): reports.replay(report) From 2660508ad68d3eca41dba009e5d3b3099fb84adf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 07:48:25 +0000 Subject: [PATCH 10/16] Test that dense widths need a checkpointed decoder Co-Authored-By: Claude Opus 5.5 (1M context) --- tests/unit/test_grouped_planner_replay.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 8868904ac..4e5d520d1 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -496,7 +496,16 @@ def test_dense_stage_on_cp_layouts_is_replayed(monkeypatch, tmp_path): @pytest.mark.parametrize( - "change", ["length", "value", "half_zero", "below_base", "no_base", "topology"] + "change", + [ + "length", + "value", + "half_zero", + "below_base", + "no_base", + "no_checkpoint", + "topology", + ], ) def test_dense_fact_validation_rejects_forged_input(change, monkeypatch, tmp_path): from test_trainer_rank_dense_memory import _dense_rank @@ -522,6 +531,10 @@ def test_dense_fact_validation_rejects_forged_input(change, monkeypatch, tmp_pat elif change == "no_base": facts["dense_base_widths"] = [0, 0] message = "invalid dense stage facts" + elif change == "no_checkpoint": + # Only a checkpointed decoder has dense widths. + facts["checkpoint_layers"] = 0 + message = "invalid dense stage facts" else: # Dense widths on a report whose recorded topology is not CP2. report["replay"]["memory_replay"]["rank"]["topology"] = [1, 1, 1, 1] From be8c68656f1c1147e90a041585abaea5d85842eb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 15:14:58 +0000 Subject: [PATCH 11/16] Refuse named-slot MoE coverage without a gradient coefficient Live coverage of a named slot needs that slot's own checkpoint coefficient (_moe_recompute_covered_for), which capture freezes as the group's gradient terms. Replay accepted moe_covered=True beside a zero coefficient and priced one covered input gradient instead of the per-boundary charge live keeps there. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 8 ++++++++ tests/unit/test_grouped_planner_replay.py | 17 +++++++++++++++++ 2 files changed, 25 insertions(+) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index ea8293399..df656e81d 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -509,6 +509,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: integer(terms[2]) if terms[2] > terms[0]: raise ValueError("invalid MoE terms") + # A named slot (only gradient groups record its adapter) is covered + # live only with its own checkpoint coefficient. + if ( + group["moe_covered"] + and group["adapter"] is not None + and not group["gradient"][0] + ): + raise ValueError("invalid MoE recompute coverage") layout = group["layout"] if layout is not None: fields(layout, {"attention_rows", "gdn_rows", "attention_retained"}) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 4e5d520d1..0f52b76cf 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -341,6 +341,23 @@ def test_recomputed_layer_fact_validation_rejects_forged_input(change, layer, tm reports.replay(report) +def test_named_slot_coverage_needs_its_gradient_coefficient(layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + group = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ][1] + assert group["adapter"] is not None and group["moe_covered"] + # A named slot without a checkpoint coefficient (an unsupported owner) + # keeps the per-boundary gradient live; coverage there would discount it. + group["gradient"] = [0, [], 0] + group["moe_covered"] = False + uncovered = reports.replay(report)["estimates"][0]["required_bytes"] + assert uncovered > 0 + group["moe_covered"] = True + with pytest.raises(ValueError, match="invalid MoE recompute coverage"): + reports.replay(report) + + @pytest.mark.parametrize("value", [None, 1.5, -1]) def test_replay_refuses_invalid_recorded_geometry(value, tmp_path): rank = head_rank() From e71ffe73d2a686a5a83108460e76aabe06415be6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 15:43:15 +0000 Subject: [PATCH 12/16] Tie replayed MoE coverage to the coefficients live records Live freezes an unnamed gradient group's terms from the constructor coefficient, covers a group only with a positive coefficient of its own, and records no MoE coefficient above TP1. Replay accepted an unnamed group with zeroed terms, a named slot relabelled unnamed, and MoE facts on a TP2 report, each pricing coverage or MoE workspace live never does. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 21 +++++++---- tests/unit/test_grouped_planner_replay.py | 36 ++++++++++++++++++- .../unit/test_planner_runtime_fact_guards.py | 3 +- 3 files changed, 52 insertions(+), 8 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index df656e81d..6e64cb976 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -509,14 +509,17 @@ def integer(value: Any, *, minimum: int = 0) -> None: integer(terms[2]) if terms[2] > terms[0]: raise ValueError("invalid MoE terms") - # A named slot (only gradient groups record its adapter) is covered - # live only with its own checkpoint coefficient. + # Only gradient groups record a named slot's adapter. Live freezes an + # unnamed one's gradient terms from the constructor coefficient, and + # covers a group only with its own positive coefficient. if ( - group["moe_covered"] - and group["adapter"] is not None - and not group["gradient"][0] + group["grad"] + and group["adapter"] is None + and group["gradient"][0] != facts["checkpoint_moe_bytes_per_token"] ): - raise ValueError("invalid MoE recompute coverage") + raise ValueError("invalid MoE terms") + if group["moe_covered"] and not group["gradient"][0]: + raise ValueError("MoE coverage without a checkpoint coefficient") layout = group["layout"] if layout is not None: fields(layout, {"attention_rows", "gdn_rows", "attention_retained"}) @@ -872,6 +875,12 @@ def runtime_arguments( # Live dense widths exist only at TP1/CP2/PP1 (_dense_mlp_widths). if any(facts["dense_base_widths"]) and self._topology_key()[1:] != (1, 2, 1): raise ValueError("dense stage facts disagree with the recorded topology") + # Live MoE coefficients are all 0 above TP1 (_moe_output_bytes_per_token). + if self._topology_key()[1] != 1 and ( + facts["checkpoint_moe_bytes_per_token"] + or any(group[key][0] for group in groups for key in ("forward", "gradient")) + ): + raise ValueError("MoE facts disagree with the recorded topology") layouts = None if groups[0]["layout"] is not None: # Live layout pricing is CP2 at TP1/PP1 only (_layout_pricing_supported). diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 0f52b76cf..7bb3b709f 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -354,7 +354,41 @@ def test_named_slot_coverage_needs_its_gradient_coefficient(layer, tmp_path): uncovered = reports.replay(report)["estimates"][0]["required_bytes"] assert uncovered > 0 group["moe_covered"] = True - with pytest.raises(ValueError, match="invalid MoE recompute coverage"): + with pytest.raises(ValueError, match="MoE coverage without a checkpoint"): + reports.replay(report) + # Relabelled unnamed, its terms must be the constructor's. + group["adapter"] = None + with pytest.raises(ValueError, match="invalid MoE terms"): + reports.replay(report) + + +@pytest.mark.parametrize("covered", [True, False]) +def test_unnamed_gradient_terms_are_the_constructor_coefficient(covered, 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 + ) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + group = facts["groups"][0] + assert group["adapter"] is None and group["moe_covered"] + assert group["gradient"][0] == facts["checkpoint_moe_bytes_per_token"] > 0 + group["moe_covered"] = covered + group["gradient"] = [0, [], 0] + with pytest.raises(ValueError, match="invalid MoE terms"): + reports.replay(report) + + +def test_moe_facts_need_the_recorded_tp1(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 + ) + memory = report["replay"]["memory_replay"] + assert memory["estimates"][0]["runtime_facts"]["checkpoint_moe_bytes_per_token"] + memory["rank"]["topology"][1] = 2 + with pytest.raises(ValueError, match="MoE facts disagree with the recorded"): reports.replay(report) diff --git a/tests/unit/test_planner_runtime_fact_guards.py b/tests/unit/test_planner_runtime_fact_guards.py index 9b2a6265f..f25dc8c4a 100644 --- a/tests/unit/test_planner_runtime_fact_guards.py +++ b/tests/unit/test_planner_runtime_fact_guards.py @@ -56,7 +56,8 @@ 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, 0] + coefficient = facts["checkpoint_moe_bytes_per_token"] + group["forward"] = [coefficient, [[0, 0]] * 4096, 0] group["gradient"] = deepcopy(group["forward"]) facts["groups"] = [deepcopy(group) for _ in range(8)] elif inventory == "slots": From 5140b2d1051897301c1ee3c07520bf482e53ddfc Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 16:06:03 +0000 Subject: [PATCH 13/16] Refuse replayed MoE facts live never records Live keeps no stages beside a zero MoE coefficient, covers a group only when the constructor coefficient is positive too, records no MoE coefficient on a dense rank or above TP1 (the rank-level forward fields included), and freezes an unnamed gradient group's forward terms from the constructor's. Replay accepted each of these forged, and several changed the priced stage or coverage. The fact-budget test's earlier edit was inert and is reverted. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 30 ++++++++-- tests/unit/test_grouped_planner_replay.py | 55 +++++++++++++++++++ .../unit/test_planner_runtime_fact_guards.py | 3 +- 3 files changed, 80 insertions(+), 8 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 6e64cb976..c990e19f9 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -507,7 +507,8 @@ def integer(value: Any, *, minimum: int = 0) -> None: integer(value) # The shared expert's part of the coefficient (live invariant). integer(terms[2]) - if terms[2] > terms[0]: + # Live keeps no stage or shared part beside a zero coefficient. + if terms[2] > terms[0] or (terms[1] and not terms[0]): raise ValueError("invalid MoE terms") # Only gradient groups record a named slot's adapter. Live freezes an # unnamed one's gradient terms from the constructor coefficient, and @@ -518,7 +519,9 @@ def integer(value: Any, *, minimum: int = 0) -> None: and group["gradient"][0] != facts["checkpoint_moe_bytes_per_token"] ): raise ValueError("invalid MoE terms") - if group["moe_covered"] and not group["gradient"][0]: + if group["moe_covered"] and not ( + group["gradient"][0] and facts["checkpoint_moe_bytes_per_token"] + ): raise ValueError("MoE coverage without a checkpoint coefficient") layout = group["layout"] if layout is not None: @@ -875,12 +878,27 @@ def runtime_arguments( # Live dense widths exist only at TP1/CP2/PP1 (_dense_mlp_widths). if any(facts["dense_base_widths"]) and self._topology_key()[1:] != (1, 2, 1): raise ValueError("dense stage facts disagree with the recorded topology") - # Live MoE coefficients are all 0 above TP1 (_moe_output_bytes_per_token). - if self._topology_key()[1] != 1 and ( - facts["checkpoint_moe_bytes_per_token"] + # Live MoE coefficients are all 0 above TP1 (_moe_output_bytes_per_token) + # and on a dense rank, which has no MoE layers (_dense_mlp_widths). + if (self._topology_key()[1] != 1 or any(facts["dense_base_widths"])) and ( + self._moe_output_bytes_per_token + or self._moe_forward_stages + or facts["checkpoint_moe_bytes_per_token"] or any(group[key][0] for group in groups for key in ("forward", "gradient")) ): - raise ValueError("MoE facts disagree with the recorded topology") + raise ValueError("MoE facts disagree with the recorded rank") + # An unnamed gradient group's forward terms are the constructor's. + for group in groups: + if ( + group["grad"] + and group["adapter"] is None + and ( + group["forward"][0] != self._moe_output_bytes_per_token + or tuple(map(tuple, group["forward"][1])) + != self._moe_forward_stages + ) + ): + raise ValueError("invalid MoE terms") layouts = None if groups[0]["layout"] is not None: # Live layout pricing is CP2 at TP1/PP1 only (_layout_pricing_supported). diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 7bb3b709f..35c86dfff 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -379,6 +379,24 @@ def test_unnamed_gradient_terms_are_the_constructor_coefficient(covered, tmp_pat reports.replay(report) +@pytest.mark.parametrize("change", ["forward_stage", "gradient_stage", "constructor"]) +def test_zero_moe_coefficients_carry_no_stage_or_coverage(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + named = facts["groups"][1] + assert named["adapter"] is not None and named["moe_covered"] + if change == "constructor": + # Live coverage needs the constructor's coefficient too. + facts["checkpoint_moe_bytes_per_token"] = 0 + message = "MoE coverage without a checkpoint" + else: + named["moe_covered"] = False + named[change.split("_")[0]] = [0, [[0, 2**40]], 0] + message = "invalid MoE terms" + with pytest.raises(ValueError, match=message): + reports.replay(report) + + def test_moe_facts_need_the_recorded_tp1(tmp_path): rank = head_rank() rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) @@ -392,6 +410,43 @@ def test_moe_facts_need_the_recorded_tp1(tmp_path): reports.replay(report) +@pytest.mark.parametrize("fact", ["rank", "checkpoint"]) +def test_dense_ranks_record_no_moe_facts(fact, monkeypatch, tmp_path): + from test_trainer_rank_dense_memory import _dense_rank + + rank = cp2(monkeypatch, _dense_rank()) + rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) + report, _ = emitted( + rank, rank._plan_flat_forward([request(512, grad=True)]), tmp_path + ) + memory = report["replay"]["memory_replay"] + facts = memory["estimates"][0]["runtime_facts"] + assert any(facts["dense_base_widths"]) + if fact == "rank": + memory["rank"]["moe_output_bytes_per_token"] = 1 + else: + # Consistent with itself (unnamed terms equal the coefficient). + facts["checkpoint_moe_bytes_per_token"] = 1 + facts["groups"][0]["gradient"] = [1, [], 0] + with pytest.raises(ValueError, match="MoE facts disagree with the recorded"): + reports.replay(report) + + +def test_unnamed_forward_terms_are_the_constructor_coefficient(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 + ) + group = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ][0] + assert group["adapter"] is None and group["forward"][0] + group["forward"] = [0, [], 0] + with pytest.raises(ValueError, match="invalid MoE terms"): + reports.replay(report) + + @pytest.mark.parametrize("value", [None, 1.5, -1]) def test_replay_refuses_invalid_recorded_geometry(value, tmp_path): rank = head_rank() diff --git a/tests/unit/test_planner_runtime_fact_guards.py b/tests/unit/test_planner_runtime_fact_guards.py index f25dc8c4a..9b2a6265f 100644 --- a/tests/unit/test_planner_runtime_fact_guards.py +++ b/tests/unit/test_planner_runtime_fact_guards.py @@ -56,8 +56,7 @@ def test_cumulative_fact_budget_precedes_json_encoding(inventory, monkeypatch): group = facts["groups"][0] group["gdn"] = None if inventory == "stages": - coefficient = facts["checkpoint_moe_bytes_per_token"] - group["forward"] = [coefficient, [[0, 0]] * 4096, 0] + group["forward"] = [0, [[0, 0]] * 4096, 0] group["gradient"] = deepcopy(group["forward"]) facts["groups"] = [deepcopy(group) for _ in range(8)] elif inventory == "slots": From 0b830be935aa382a0372fb86b555d6062cdef7b6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 16:41:15 +0000 Subject: [PATCH 14/16] Keep replayed MoE terms consistent across modes and groups One walk prices both modes, so live's checkpoint coefficient is never below the forward one, and the same storage gates both modes' stages. Unnamed gradient groups all read the constructor's coverage flag, which a named slot's coverage also needs. Replay accepted each of these forged. The named-slot test's uncovered baseline now zeroes both modes, as an unsupported slot does live; two round-4 tests now isolate the invariant they name. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 15 ++++++++ tests/unit/test_grouped_planner_replay.py | 44 ++++++++++++++++++++--- 2 files changed, 55 insertions(+), 4 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index c990e19f9..89be5aed6 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -523,6 +523,13 @@ def integer(value: Any, *, minimum: int = 0) -> None: group["gradient"][0] and facts["checkpoint_moe_bytes_per_token"] ): raise ValueError("MoE coverage without a checkpoint coefficient") + # One walk prices both modes: the checkpoint coefficient is at least + # the forward one, and the same storage gates both modes' stages. + forward, gradient = group["forward"], group["gradient"] + if gradient[0] < forward[0] or ( + forward[0] and bool(forward[1]) != bool(gradient[1]) + ): + raise ValueError("invalid MoE terms") layout = group["layout"] if layout is not None: fields(layout, {"attention_rows", "gdn_rows", "attention_retained"}) @@ -597,6 +604,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: } if len(kinds) > 1: raise ValueError("invalid adapter gradient facts") + # Unnamed gradient groups all read the constructor's coverage flag, and a + # named slot is covered only where that flag is set. + unnamed = {g["moe_covered"] for g in groups if g["grad"] and g["adapter"] is None} + if len(unnamed) > 1 or ( + False in unnamed + and any(g["moe_covered"] for g in groups if g["adapter"] is not None) + ): + raise ValueError("inconsistent MoE recompute coverage") # Live layouts cover every group of a plan on the same ranks, or none. layouts = [group["layout"] for group in groups] laid_out = any(layout is not None for layout in layouts) diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index 35c86dfff..f1431299e 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -347,8 +347,9 @@ def test_named_slot_coverage_needs_its_gradient_coefficient(layer, tmp_path): "groups" ][1] assert group["adapter"] is not None and group["moe_covered"] - # A named slot without a checkpoint coefficient (an unsupported owner) - # keeps the per-boundary gradient live; coverage there would discount it. + # A named slot without a coefficient (an unsupported owner, zero in both + # modes) keeps the per-boundary gradient live; coverage would discount it. + group["forward"] = [0, [], 0] group["gradient"] = [0, [], 0] group["moe_covered"] = False uncovered = reports.replay(report)["estimates"][0]["required_bytes"] @@ -386,8 +387,12 @@ def test_zero_moe_coefficients_carry_no_stage_or_coverage(change, layer, tmp_pat named = facts["groups"][1] assert named["adapter"] is not None and named["moe_covered"] if change == "constructor": - # Live coverage needs the constructor's coefficient too. + # Live coverage needs the constructor's coefficient too (zero it + # whole, forward included, so only coverage is inconsistent). facts["checkpoint_moe_bytes_per_token"] = 0 + rank_fields = report["replay"]["memory_replay"]["rank"] + rank_fields["moe_output_bytes_per_token"] = 0 + rank_fields["moe_forward_stages"] = [] message = "MoE coverage without a checkpoint" else: named["moe_covered"] = False @@ -397,6 +402,37 @@ def test_zero_moe_coefficients_carry_no_stage_or_coverage(change, layer, tmp_pat reports.replay(report) +@pytest.mark.parametrize("change", ["below_forward", "stages", "coverage"]) +def test_moe_terms_agree_across_modes_and_groups(change, layer, tmp_path): + report, _, _ = adapter_report(layer, tmp_path) + groups = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ + "groups" + ] + named = groups[1] + assert named["adapter"] is not None and named["moe_covered"] + assert named["forward"][0] <= named["gradient"][0] and named["gradient"][1] + if change == "below_forward": + named["gradient"] = [1, [], 0] + message = "invalid MoE terms" + elif change == "stages": + named["forward"][1] = [] + message = "invalid MoE terms" + else: + from art.trainer_rank import _planner_replay + + # An uncovered unnamed gradient group beside a covered named slot; + # validation refuses it before the recorded arguments are compared. + facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] + groups[0]["grad"] = True + groups[0]["gradient"] = deepcopy(named["gradient"]) + groups[0]["gradient"][0] = facts["checkpoint_moe_bytes_per_token"] + with pytest.raises(ValueError, match="inconsistent MoE recompute coverage"): + _planner_replay.validate(facts) + return + with pytest.raises(ValueError, match=message): + reports.replay(report) + + def test_moe_facts_need_the_recorded_tp1(tmp_path): rank = head_rank() rank._planner_reporter = reports.Reporter(0, spool_dir=tmp_path) @@ -423,7 +459,7 @@ def test_dense_ranks_record_no_moe_facts(fact, monkeypatch, tmp_path): facts = memory["estimates"][0]["runtime_facts"] assert any(facts["dense_base_widths"]) if fact == "rank": - memory["rank"]["moe_output_bytes_per_token"] = 1 + memory["rank"]["moe_forward_stages"] = [[0, 1]] else: # Consistent with itself (unnamed terms equal the coefficient). facts["checkpoint_moe_bytes_per_token"] = 1 From 85423ed689a2bfb012477198ede5d0b1823a7424 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 16:57:49 +0000 Subject: [PATCH 15/16] Pin replayed MoE terms to one walk in both modes Live's two modes share one walk: its coefficients are zero together, its stages exist together, and the checkpoint pass adds a stage per converted FC2. Replay accepted a zero forward coefficient beside a positive gradient one, and a gradient stage list no longer than the forward one. Tests isolate the coefficient-order case and cover unnamed groups that disagree on coverage. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_planner_replay.py | 11 ++++--- tests/unit/test_grouped_planner_replay.py | 37 +++++++++++++++++++---- 2 files changed, 38 insertions(+), 10 deletions(-) diff --git a/src/art/trainer_rank/_planner_replay.py b/src/art/trainer_rank/_planner_replay.py index 89be5aed6..f08291503 100644 --- a/src/art/trainer_rank/_planner_replay.py +++ b/src/art/trainer_rank/_planner_replay.py @@ -523,11 +523,14 @@ def integer(value: Any, *, minimum: int = 0) -> None: group["gradient"][0] and facts["checkpoint_moe_bytes_per_token"] ): raise ValueError("MoE coverage without a checkpoint coefficient") - # One walk prices both modes: the checkpoint coefficient is at least - # the forward one, and the same storage gates both modes' stages. + # One walk prices both modes under the same gates: its checkpoint + # pass only raises the coefficient and adds a stage per converted FC2. forward, gradient = group["forward"], group["gradient"] - if gradient[0] < forward[0] or ( - forward[0] and bool(forward[1]) != bool(gradient[1]) + if ( + gradient[0] < forward[0] + or bool(gradient[0]) != bool(forward[0]) + or bool(gradient[1]) != bool(forward[1]) + or (forward[1] and len(gradient[1]) <= len(forward[1])) ): raise ValueError("invalid MoE terms") layout = group["layout"] diff --git a/tests/unit/test_grouped_planner_replay.py b/tests/unit/test_grouped_planner_replay.py index f1431299e..8303b7a47 100644 --- a/tests/unit/test_grouped_planner_replay.py +++ b/tests/unit/test_grouped_planner_replay.py @@ -402,7 +402,17 @@ def test_zero_moe_coefficients_carry_no_stage_or_coverage(change, layer, tmp_pat reports.replay(report) -@pytest.mark.parametrize("change", ["below_forward", "stages", "coverage"]) +@pytest.mark.parametrize( + "change", + [ + "below_forward", + "stages", + "zero_forward", + "stage_count", + "coverage", + "unnamed_coverage", + ], +) def test_moe_terms_agree_across_modes_and_groups(change, layer, tmp_path): report, _, _ = adapter_report(layer, tmp_path) groups = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"][ @@ -410,22 +420,35 @@ def test_moe_terms_agree_across_modes_and_groups(change, layer, tmp_path): ] named = groups[1] assert named["adapter"] is not None and named["moe_covered"] - assert named["forward"][0] <= named["gradient"][0] and named["gradient"][1] + assert named["forward"][0] <= named["gradient"][0] + assert 0 < len(named["forward"][1]) < len(named["gradient"][1]) if change == "below_forward": - named["gradient"] = [1, [], 0] + # Stages kept, so only the coefficient order is wrong. + named["gradient"][0] = named["forward"][0] - 1 message = "invalid MoE terms" elif change == "stages": named["forward"][1] = [] message = "invalid MoE terms" + elif change == "zero_forward": + named["forward"] = [0, [], 0] + message = "invalid MoE terms" + elif change == "stage_count": + named["gradient"][1] = named["gradient"][1][: len(named["forward"][1])] + message = "invalid MoE terms" else: from art.trainer_rank import _planner_replay - # An uncovered unnamed gradient group beside a covered named slot; - # validation refuses it before the recorded arguments are compared. + # Validation refuses these before the recorded arguments are compared: + # an uncovered unnamed gradient group beside a covered named slot, or + # two unnamed gradient groups that disagree. facts = report["replay"]["memory_replay"]["estimates"][0]["runtime_facts"] groups[0]["grad"] = True groups[0]["gradient"] = deepcopy(named["gradient"]) groups[0]["gradient"][0] = facts["checkpoint_moe_bytes_per_token"] + if change == "unnamed_coverage": + named["moe_covered"] = False + groups.append(deepcopy(groups[0])) + groups[-1]["moe_covered"] = True with pytest.raises(ValueError, match="inconsistent MoE recompute coverage"): _planner_replay.validate(facts) return @@ -461,8 +484,10 @@ def test_dense_ranks_record_no_moe_facts(fact, monkeypatch, tmp_path): if fact == "rank": memory["rank"]["moe_forward_stages"] = [[0, 1]] else: - # Consistent with itself (unnamed terms equal the coefficient). + # Consistent with itself (unnamed terms equal the coefficient in + # both modes). facts["checkpoint_moe_bytes_per_token"] = 1 + facts["groups"][0]["forward"] = [1, [], 0] facts["groups"][0]["gradient"] = [1, [], 0] with pytest.raises(ValueError, match="MoE facts disagree with the recorded"): reports.replay(report) From 61a3a20a65927abdf7560f87a8e2dbe87712a56d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 12:14:00 -0600 Subject: [PATCH 16/16] Allow a configured free Kubernetes context for required GPU CI (#1045) (cherry picked from commit 8428ef8f0188a6a9cb08e632085ea4929d89f665) --- .github/workflows/trainer-rank-gpu.yml | 4 ++-- CONTRIBUTING.md | 6 +++++- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/.github/workflows/trainer-rank-gpu.yml b/.github/workflows/trainer-rank-gpu.yml index 088b53122..6bc5849be 100644 --- a/.github/workflows/trainer-rank-gpu.yml +++ b/.github/workflows/trainer-rank-gpu.yml @@ -146,7 +146,7 @@ jobs: timeout-minutes: 45 environment: trainer-rank-gpu-validation env: - SKY_INFRA: k8s/${{ inputs.context || 'cks-wb3' }} + SKY_INFRA: k8s/${{ inputs.context || vars.TRAINER_RANK_GPU_CONTEXT || 'cks-wb3' }} steps: - uses: actions/checkout@v4 with: @@ -172,7 +172,7 @@ jobs: - name: Configure Kubernetes env: - KUBECONFIG_VALUE: ${{ secrets[inputs.context == 'ext-collab2' && 'GPU_IMAGE_KUBECONFIG' || 'CKS_WB3_KUBECONFIG'] }} + KUBECONFIG_VALUE: ${{ secrets[env.SKY_INFRA == 'k8s/ext-collab2' && 'GPU_IMAGE_KUBECONFIG' || 'CKS_WB3_KUBECONFIG'] }} run: | set -euo pipefail case "${SKY_INFRA}" in diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 6f9dcdb9f..87b3546bb 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -40,13 +40,17 @@ These checks are automatically run in CI for all pull requests. If your PR fails ### TrainerRank GPU Validation GPU validation uses two H200s on free `cks-wb3` Kubernetes infrastructure by default. +Set the repository Actions variable `TRAINER_RANK_GPU_CONTEXT` to `ext-collab2` +to place required PR GPU jobs there instead; unset it to restore the default. +This selects infrastructure without reserving capacity or changing validation. To select `ext-collab2` for a manual validation of a branch: ```bash gh workflow run trainer-rank-gpu.yml --ref BRANCH -f context=ext-collab2 ``` -These are the only supported contexts; there is no automatic fallback. The default +The manual input takes precedence over the repository variable. These are the +only supported contexts; there is no automatic fallback. The default uses `CKS_WB3_KUBECONFIG`; `ext-collab2` uses the existing `GPU_IMAGE_KUBECONFIG` secret and requires its `skypilot-workload` service account. The selected infrastructure is recorded in the job's owner and result receipts. Source checks, time limits,