From 29024e3d6e2f552a3fa44ec6cd75468619cf9dc5 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 08:20:42 +0000 Subject: [PATCH 1/7] 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 2/7] 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 3/7] 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 4/7] 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 5/7] 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 6/7] 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 7/7] 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())